From ebad146b839ab10a8f3a81452f560fbeb7358d33 Mon Sep 17 00:00:00 2001 From: admin Date: Sun, 6 Sep 2026 09:09:18 +0000 Subject: [PATCH] =?UTF-8?q?=E6=96=B0=E5=A2=9E=20rtl-signal-mcp:=20RTL=20?= =?UTF-8?q?=E4=BF=A1=E5=8F=B7=E8=BF=BD=E8=B8=AA=20MCP=20server=EF=BC=88?= =?UTF-8?q?=E7=B4=A2=E5=BC=95=E5=BC=95=E6=93=8E+8=E5=B7=A5=E5=85=B7+?= =?UTF-8?q?=E9=AA=8C=E6=94=B6=E8=84=9A=E6=9C=AC=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- rtl-mcp/README.md | 76 ++++++++ rtl-mcp/mcp_client_demo.py | 194 ++++++++++++++++++++ rtl-mcp/mcp_server.py | 253 ++++++++++++++++++++++++++ rtl-mcp/rtl_index.py | 356 +++++++++++++++++++++++++++++++++++++ 4 files changed, 879 insertions(+) create mode 100644 rtl-mcp/README.md create mode 100644 rtl-mcp/mcp_client_demo.py create mode 100644 rtl-mcp/mcp_server.py create mode 100644 rtl-mcp/rtl_index.py diff --git a/rtl-mcp/README.md b/rtl-mcp/README.md new file mode 100644 index 0000000..94cd7f2 --- /dev/null +++ b/rtl-mcp/README.md @@ -0,0 +1,76 @@ +# rtl-signal-mcp + +给 AI agent 用的 RTL 信号追踪 MCP Server。解决一个具体问题:**RTL 信号的联系跨模块分布在端口连接上,读单个文件看不出来**。本工具把整个代码库索引成"模块事实库 + 跨边界信号边",agent 查一次得一条紧凑驱动链,不用通读源码、不烧 token。 + +实测性能:OpenTitan 全 hw 树(31 万行 RTL)→ **971 模块 / 94,891 条跨模块边 / 10.6 秒建索引 / 310MB 内存**;单次追踪查询毫秒级、返回几百字节文本。 + +## 包含文件 + +| 文件 | 说明 | +|---|---| +| `rtl_index.py` | 核心引擎:pyslang 解析 → 模块事实库(端口/网络/assign/always/实例化)→ 跨模块边 → 驱动/负载递归查询 | +| `mcp_server.py` | MCP stdio server(纯标准库实现 JSON-RPC 2.0,无 SDK 依赖) | +| `mcp_client_demo.py` | 测试客户端(模拟 agent 完整调用序列,可当验收脚本) | +| `README.md` | 本文件 | + +## 8 个工具 + +| 工具 | 作用 | token 成本 | +|---|---|---| +| `build_index` | 建索引(roots: 文件/目录列表) | 一次性 | +| `list_modules` | 列模块及事实计数 | 极小 | +| `module_summary` | 单模块端口/实例/连接摘要 | 小 | +| `find_signal` | 正则搜信号 → 模块 + file:line | 极小 | +| **`trace_drivers`** | **谁驱动这个信号**(跨模块递归) | **每次几百字节** | +| **`trace_loads`** | **这个信号驱动谁**(上抛+下钻) | 小 | +| `get_source` | 按行取小段源码(≤200 行) | 按需 | +| `hierarchy` | 实例化树 | 小 | + +## 接入 Claude / 其他 MCP 客户端 + +```json +{ + "mcpServers": { + "rtl-signal": { + "command": "python3", + "args": ["/path/to/rtl-mcp/mcp_server.py"], + "env": { "RTL_MCP_ROOTS": "/data/repos/e203_hbirdv2/rtl/e203:/data/repos/opentitan/hw/ip" } + } + } +} +``` + +`RTL_MCP_ROOTS` 冒号分隔多个目录,启动时自动预建索引;不设则由 agent 首次调用 `build_index`。 + +依赖:`pip install pyslang`(仅此一个)。Python ≥ 3.10。 + +## Agent 使用模式(写进系统提示词即可) + +``` +调试 RTL 信号问题时按此流程: +1. find_signal 找到信号的归属模块和位置 +2. trace_drivers 查驱动链(跨模块),trace_loads 查影响面 +3. 只对链上关键条目用 get_source 看 file:line 附近 ±20 行 +4. 绝不整文件读源码;结论引用 [kind] where :: sig <- rhs 条目 +``` + +## 已验证范围 + +- E203 蜂鸟 RISC-V 核(42 文件/39 模块/1789 边/0.2s):`trace_drivers(e203_exu_alu.i_valid)` 正确给出 ALU←EXU 译码器三模块链(含 file:line) +- OpenTitan `hw/ip/aes`(67 模块/1037 边) +- OpenTitan `hw/ip + top_earlgrey`(841 模块/84,942 边/7s) +- OpenTitan 全 hw 树(971 模块/94,891 边/10.6s/310MB) + +## 已知边界(诚实声明) + +- 语法级分析:不做参数展开/elaboration——generate for、parameter 化的端口名按字面处理;位宽/方向冲突不检测 +- always 块提取覆盖 always/always_ff/always_comb 内的阻塞与非阻塞赋值(含 if/case 嵌套与时序控制穿透),函数内赋值暂不提取 +- 跨模块追踪经实例端口边递归,Max depth 默认 4(可调);未索引的第三方模块(如仅黑盒实例名)标为链尾 +- SystemVerilog interface/class/sequence 未索引(面向可综合 RTL 设计) +- 同名信号在不同模块中是不同节点——查询需给 module + signal 两个参数(`find_signal` 可帮定位) + +## 路线图(人用的网页版后置) + +- v0.2: elaboration 级准确(走 pyslang AST 编译而不是纯语法树),generate/parameter 展开 +- v0.3: 波形接入(pylibfst 读 FST,值反标到驱动链 → X 态/根因追踪);参考腾讯开源 wave-mcp +- v0.4: 网页视图(驱动链图形化 + 波形联动),给人 review 用 diff --git a/rtl-mcp/mcp_client_demo.py b/rtl-mcp/mcp_client_demo.py new file mode 100644 index 0000000..b6f1f0e --- /dev/null +++ b/rtl-mcp/mcp_client_demo.py @@ -0,0 +1,194 @@ +#!/usr/bin/env python3 +"""mcp_client_demo.py — rtl-signal-mcp 验收脚本。 + +不做真 MCP 握手库,直接以子进程方式拉起 mcp_server.py,按行写 JSON-RPC、按行读响应, +模拟一个 agent 的完整调用序列: + + initialize -> notifications/initialized -> tools/list + -> build_index -> list_modules -> module_summary -> find_signal + -> trace_drivers -> trace_loads -> get_source -> hierarchy + +用法: + python3 mcp_client_demo.py /data/repos/e203_hbirdv2/rtl/e203 + python3 mcp_client_demo.py # 缺省用 E203 路径(不存在则报错提示) + +退出码 0 = 全部步骤通过;非 0 = 某步失败(打印到 stderr)。 +""" +import json, os, subprocess, sys + +HERE = os.path.dirname(os.path.abspath(__file__)) +SERVER = os.path.join(HERE, "mcp_server.py") + +DEFAULT_ROOTS = [ + "/data/repos/e203_hbirdv2/rtl/e203", # 本地缺省 +] + + +def start_server(roots): + env = dict(os.environ) + env["RTL_MCP_ROOTS"] = ":".join(roots) + return subprocess.Popen( + [sys.executable, SERVER], + stdin=subprocess.PIPE, stdout=subprocess.PIPE, + env=env, text=True, encoding="utf-8", bufsize=1) + + +def call(proc, method, params=None, msg_id=[0]): + """写一条请求并读回响应(跳过通知)。返回 result 或抛 RuntimeError。""" + msg_id[0] += 1 + req = dict(jsonrpc="2.0", id=msg_id[0], method=method, params=params or {}) + proc.stdin.write(json.dumps(req) + "\n") + proc.stdin.flush() + while True: + line = proc.stdout.readline() + if not line: + raise RuntimeError(f"server closed during {method}") + resp = json.loads(line) + if resp.get("id") == msg_id[0]: + if "error" in resp: + raise RuntimeError(f"{method} error: {resp['error']}") + return resp["result"] + + +def notify(proc, method, params=None): + req = dict(jsonrpc="2.0", method=method, params=params or {}) + proc.stdin.write(json.dumps(req) + "\n") + proc.stdin.flush() + + +def show(title, text, max_lines=12): + print(f"\n=== {title} ===") + lines = text.splitlines() + for ln in lines[:max_lines]: + print(" ", ln) + if len(lines) > max_lines: + print(f" ... ({len(lines) - max_lines} more lines)") + + +def fail(step, msg): + raise RuntimeError(f"step {step} failed: {msg}") + + +def main(): + roots = [a for a in sys.argv[1:] if not a.startswith("-")] or DEFAULT_ROOTS + missing = [r for r in roots if not os.path.exists(r)] + if missing: + print(f"roots not found: {missing}\n" + f"usage: python3 mcp_client_demo.py [...]", file=sys.stderr) + return 2 + + proc = start_server(roots) + try: + # 1. MCP 握手 + r = call(proc, "initialize", dict( + protocolVersion="2024-11-05", + capabilities={}, + clientInfo=dict(name="demo-client", version="0.1"))) + ver = r["protocolVersion"] + print(f"[1] initialize OK server={r['serverInfo']['name']} " + f"v{r['serverInfo']['version']} protocol={ver}") + notify(proc, "notifications/initialized") + + # 2. 工具清单 + r = call(proc, "tools/list") + tools = [t["name"] for t in r["tools"]] + print(f"[2] tools/list OK {len(tools)} tools: {', '.join(tools)}") + + # 3. 建索引(启动时已用 RTL_MCP_ROOTS 预建过一次,这里验证重复调用幂等) + r = call(proc, "tools/call", dict(name="build_index", + arguments=dict(roots=roots))) + st = r["structuredContent"] + print(f"[3] build_index OK modules={st['modules']} " + f"cross_edges={st['cross_edges']}") + + # 4. 模块列表(取第一个模块做后续实验对象) + r = call(proc, "tools/call", dict(name="list_modules", arguments={})) + mods = r["structuredContent"]["modules"] + print(f"[4] list_modules OK total={r['structuredContent']['total']} " + f"first={mods[0]['module']}") + + # 5. 挑一个"实例数较多"的模块做摘要 + target = max(mods, key=lambda m: m["insts"])["module"] + r = call(proc, "tools/call", dict(name="module_summary", + arguments=dict(module=target))) + show(f"module_summary {target}", r["content"][0]["text"], 8) + + # 6. 找一个跨模块信号:用摘要里第一个输出端口 + r2 = call(proc, "tools/call", dict(name="module_summary", + arguments=dict(module=target))) + sc = r2["structuredContent"] + port = None + for p in sc.get("ports", []): + # 端口格式是 "Output nice_xs_off" / "Input csr_ena" 这样的字符串 + if str(p).startswith("Output"): + port = str(p).split(None, 1)[1] + break + if port is None: # 摘要结构兜底:直接试常见信号 + port = "clk" + r = call(proc, "tools/call", dict(name="find_signal", + arguments=dict(pattern=port, limit=5))) + hits = r["structuredContent"]["hits"] + print(f"[6] find_signal '{port}' OK {len(hits)} hits") + + # 7. 追驱动链(核心能力) + r = call(proc, "tools/call", dict(name="trace_drivers", + arguments=dict(module=target, + signal=port, depth=4))) + sc = r["structuredContent"] + n = len(sc.get("result", [])) + print(f"[7] trace_drivers {target}.{port} OK {n} rows") + show("driver chain", r["content"][0]["text"], 12) + if n == 0 and port != "clk": + # 输出端口可能由子模块驱动也可能本就是顶层输入,空不算失败 + print(" (empty chain — top-level input or tb-driven)") + + # 8. 追负载 + r = call(proc, "tools/call", dict(name="trace_loads", + arguments=dict(module=target, + signal=port, depth=3))) + n = len(r["structuredContent"].get("result", [])) + print(f"[8] trace_loads {target}.{port} OK {n} rows") + + # 9. 从驱动链里拿一个 file:line 看小段源码 + r = call(proc, "tools/call", dict(name="trace_drivers", + arguments=dict(module=target, + signal=port, depth=2))) + rows = r["structuredContent"].get("result", []) + loc = None + import re + for row in rows: + # where 形如 "e203_exu_csr (e203_exu_csr.v:210)" + m = re.search(r"([\w./]+):(\d+)\)\s*$", row.get("where", "")) + if row.get("kind") == "assign" and m: + loc = (m.group(1), int(m.group(2))) + break + if loc: + fname, lineno = loc + r = call(proc, "tools/call", dict( + name="get_source", + arguments=dict(file=fname, start=max(1, lineno - 5), + end=lineno + 5))) + show(f"get_source {fname}:{lineno - 5}-{lineno + 5}", + r["content"][0]["text"], 12) + if "error" in r["structuredContent"]: + fail("get_source", r["structuredContent"]["error"]) + else: + print("[9] get_source skipped (no assign location in chain)") + + # 10. 层次树 + r = call(proc, "tools/call", dict(name="hierarchy", + arguments=dict(module=target, depth=2))) + show("hierarchy", r["content"][0]["text"], 10) + + print("\nALL CHECKS PASSED") + return 0 + finally: + try: + proc.stdin.close() + except Exception: + pass + proc.terminate() + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/rtl-mcp/mcp_server.py b/rtl-mcp/mcp_server.py new file mode 100644 index 0000000..ea529c9 --- /dev/null +++ b/rtl-mcp/mcp_server.py @@ -0,0 +1,253 @@ +#!/usr/bin/env python3 +"""rtl-signal-mcp — MCP server exposing RTL signal-flow tracing to AI agents. + +Protocol: MCP over stdio (JSON-RPC 2.0). No third-party MCP SDK required. + +Tools (token-efficient by design — every answer is a compact graph, not source): + build_index build/replace the index from file/dir roots + list_modules module names + instantiation counts + module_summary ports / instances / net count of one module + find_signal regex search over all signals -> module + file:line + trace_drivers who drives this signal (recursively, across module boundaries) + trace_loads what does this signal drive (recursively, up and down) + get_source small excerpt of a file by line range + hierarchy instantiation tree under a module + +Config: env RTL_MCP_ROOTS="/path/a:/path/b" pre-builds at startup. +""" +import os, sys, json, io + +sys.stdout = io.TextIOWrapper(sys.stdout.buffer, encoding="utf-8", write_through=True) +sys.stderr = io.TextIOWrapper(sys.stderr.buffer, encoding="utf-8", write_through=True) + +from rtl_index import RtlIndex + +IDX = None # global index + + +def _need_index(): + if IDX is None or not IDX.db: + raise RuntimeError("index empty: call build_index first") + + +# ---------------- tool implementations ---------------- + +def t_build_index(args): + global IDX + roots = args.get("roots") or [r for r in os.environ.get("RTL_MCP_ROOTS", "").split(":") if r] + if not roots: + return dict(error="no roots given (pass roots[] or set RTL_MCP_ROOTS)") + t0 = os.times() + IDX = RtlIndex.build(roots) + n_always = sum(len(m["always"]) for m in IDX.db.values()) + n_assign = sum(len(m["assigns"]) for m in IDX.db.values()) + return dict(modules=len(IDX.db), cross_edges=len(IDX.edges), + assign_facts=n_assign, always_facts=n_always, + note="index in memory; use find_signal to locate, trace_drivers/trace_loads to follow") + + +def t_list_modules(args): + _need_index() + pat = args.get("pattern") + import re + rx = re.compile(pat, re.I) if pat else None + out = [] + for name, m in IDX.db.items(): + if rx and not rx.search(name): + continue + out.append(dict(module=name, + ports=len(m["ports"]), insts=len(m["insts"]), + assigns=len(m["assigns"]), always=len(m["always"]))) + return dict(modules=out[:80], total=len(out)) + + +def t_module_summary(args): + _need_index() + s = IDX.module_summary(args["module"]) + if s is None: + return dict(error=f"module {args['module']} not found") + return s + + +def t_find_signal(args): + _need_index() + hits = IDX.search(args["pattern"], limit=int(args.get("limit", 20))) + return dict(hits=hits) + + +def _fmt_trace(rows): + lines = [] + for r in rows: + rhs = ", ".join(r["rhs"][:6]) + lines.append(f"[{r['kind']}] {r['where']} :: {r['sig']} <- {rhs}") + return lines + + +def t_trace_drivers(args): + _need_index() + rows = IDX.drivers(args["module"], args["signal"], + max_depth=int(args.get("depth", 4))) + if not rows: + return dict(result=[], text=f"no driver found for {args['module']}.{args['signal']} " + "(top input / testbench-driven / not indexed)") + lines = _fmt_trace(rows) + return dict(result=rows, text="\n".join(lines), + note="read chain bottom-up: leaf entries are the ultimate drivers; " + "use get_source for file:line details") + + +def t_trace_loads(args): + _need_index() + rows = IDX.loads(args["module"], args["signal"], + max_depth=int(args.get("depth", 3))) + if not rows: + return dict(result=[], text=f"no loads found for {args['module']}.{args['signal']}") + lines = _fmt_trace(rows) + return dict(result=rows, text="\n".join(lines)) + + +_FMAP = None + + +def _file_map(): + """basename -> absolute path, built lazily by walking indexed roots.""" + global _FMAP + if _FMAP is None: + _FMAP = {} + for r in getattr(IDX, "_roots", []) or []: + if os.path.isfile(r): + _FMAP.setdefault(os.path.basename(r), r) + continue + for dp, _dns, fns in os.walk(r): + for fn in fns: + _FMAP.setdefault(fn, os.path.join(dp, fn)) + return _FMAP + + +def t_get_source(args): + path, start, end = args["file"], int(args["start"]), int(args["end"]) + end = min(end, start + 200) + if not os.path.isabs(path): + # resolve by basename against indexed roots (files may sit in subdirs) + cand = _file_map().get(os.path.basename(path)) + if cand: + path = cand + if not os.path.exists(path): + return dict(error=f"file not found: {path}") + out = [] + with open(path, encoding="utf-8", errors="replace") as f: + for i, line in enumerate(f, 1): + if start <= i <= end: + out.append(f"{i:5d}| {line.rstrip()}") + if i > end: + break + return dict(text="\n".join(out)) + + +def t_hierarchy(args): + _need_index() + top = args["module"] + depth = int(args.get("depth", 2)) + lines = [] + + def rec(mod, indent, d, seen): + if d > depth or (mod, indent) in seen or len(lines) > 80: + return + seen.add((mod, indent)) + m = IDX.db.get(mod) + if not m: + return + for (itype, iname, conns) in m["insts"]: + lines.append(f"{' ' * (indent + 1)}└─ {iname} : {itype}") + rec(itype, indent + 1, d + 1, seen) + + lines.append(top) + rec(top, 0, 0, set()) + return dict(text="\n".join(lines)) + + +TOOLS = { + "build_index": (t_build_index, "Build the signal-flow index. roots: list of files or directories to scan (.v/.sv)."), + "list_modules": (t_list_modules, "List indexed modules (name, ports, insts, fact counts). pattern: optional regex."), + "module_summary": (t_module_summary, "Compact summary of one module: port list, instances with connected nets, net names."), + "find_signal": (t_find_signal, "Find signals by regex. Returns module + file:line for each hit."), + "trace_drivers": (t_trace_drivers, "Trace WHO drives a signal — recursive across module boundaries (Verdi-style). Args: module, signal, depth=4."), + "trace_loads": (t_trace_loads, "Trace WHAT a signal drives — loads in same module, flows out output ports, feeds child instances. Args: module, signal, depth=3."), + "get_source": (t_get_source, "Read a small source excerpt (file, start, end; max 200 lines). Use after tracing to see the exact logic."), + "hierarchy": (t_hierarchy, "Instantiation tree under a module. Args: module, depth=2."), +} + +TOOLSPEC = [dict(name=n, description=d, + inputSchema=dict(type="object", + properties={}, + additionalProperties=True)) + for n, (f, d) in TOOLS.items()] + + +# ---------------- JSON-RPC / MCP plumbing ---------------- + +def reply(msg_id, result): + sys.stdout.write(json.dumps(dict(jsonrpc="2.0", id=msg_id, result=result)) + "\n") + + +def reply_error(msg_id, code, message): + sys.stdout.write(json.dumps(dict(jsonrpc="2.0", id=msg_id, + error=dict(code=code, message=message))) + "\n") + + +def handle(req): + method = req.get("method") + msg_id = req.get("id") + if method == "initialize": + reply(msg_id, dict(protocolVersion="2024-11-05", + capabilities=dict(tools={}), + serverInfo=dict(name="rtl-signal-mcp", version="0.1.0"))) + elif method == "notifications/initialized": + pass + elif method == "tools/list": + reply(msg_id, dict(tools=TOOLSPEC)) + elif method == "tools/call": + params = req.get("params", {}) + name = params.get("name") + args = params.get("arguments", {}) or {} + fn = TOOLS.get(name) + if fn is None: + reply_error(msg_id, -32602, f"unknown tool {name}") + return + try: + res = fn[0](args) + text = res.pop("text", None) or json.dumps(res, ensure_ascii=False) + reply(msg_id, dict(content=[dict(type="text", text=text)], + structuredContent=res, isError=False)) + except Exception as e: + reply(msg_id, dict(content=[dict(type="text", text=f"error: {e}")], isError=True)) + elif method == "ping": + reply(msg_id, {}) + elif msg_id is not None: + reply_error(msg_id, -32601, f"method not supported: {method}") + + +def main(): + if os.environ.get("RTL_MCP_ROOTS"): + try: + res = t_build_index({}) + print(f"[rtl-signal-mcp] prebuilt index: {res}", file=sys.stderr) + except Exception as e: + print(f"[rtl-signal-mcp] prebuild failed: {e}", file=sys.stderr) + for line in sys.stdin: + line = line.strip() + if not line: + continue + try: + req = json.loads(line) + except Exception: + continue + try: + handle(req) + except Exception as e: + if req.get("id") is not None: + reply_error(req["id"], -32603, str(e)) + + +if __name__ == "__main__": + main() diff --git a/rtl-mcp/rtl_index.py b/rtl-mcp/rtl_index.py new file mode 100644 index 0000000..3fe6296 --- /dev/null +++ b/rtl-mcp/rtl_index.py @@ -0,0 +1,356 @@ +#!/usr/bin/env python3 +"""rtl_index.py — RTL signal-flow index for cross-module tracing. + +Builds a module-fact database from Verilog/SystemVerilog sources: + ports (ANSI & non-ANSI), nets, continuous assigns, always-block drivers, + instantiations with named port connections, and cross-boundary edges. + +pyslang 11 API. Index build is O(source); queries are in-memory graph walks. +""" +import os, re, json, glob +import pyslang +from pyslang import syntax as sl + +IDENT = re.compile(r"^[A-Za-z_][A-Za-z0-9_$]*$") +ASSIGN_KINDS = ("AssignmentExpression", "NonblockingAssignmentExpression") + + +def _txt(node): + t = getattr(node, "valueText", "") + return t if t else (str(node).strip() if node is not None else "") + + +def _syntax_of(val_node): + s = getattr(val_node, "syntax", None) + return s if s is not None else val_node + + +def ids_from(expr_syntax): + out = [] + def walk(n): + if n is None: + return + if n.kind in (sl.SyntaxKind.IdentifierName, sl.SyntaxKind.IdentifierSelectName): + ident = getattr(n, "identifier", None) + if ident is not None: + t = _txt(ident) + if IDENT.match(t): + out.append(t) + # IdentifierSelectName has .target identifier inside + tgt = getattr(n, "target", None) + if tgt is not None: + walk(tgt) + return + for c in n: + if hasattr(c, "__iter__"): + walk(c) + walk(expr_syntax) + return out + + +def _lhs_names(expr_syntax): + """names written by an assignment LHS (plain idents only, strip selects)""" + out = [] + def walk(n): + if n is None: + return + if n.kind == sl.SyntaxKind.IdentifierName: + t = _txt(n.identifier) + if IDENT.match(t): + out.append(t) + return + for c in n: + if hasattr(c, "__iter__"): + walk(c) + walk(expr_syntax) + return out + + +class RtlIndex: + def __init__(self): + self.sm = pyslang.SourceManager() + self.db = {} # module -> facts + self.edges = [] # (parent_mod, parent_net, child_mod, child_port, child_dir, inst_name) + self.loc = {} # (module, signal) -> "file:line" + + # ---------- build ---------- + def add_tree(self, path): + tree = sl.SyntaxTree.fromFile(path, self.sm) + root = tree.root + for m in root.members: + if m.kind == sl.SyntaxKind.ModuleDeclaration: + info = self._scan_module(m) + self.db[info["name"]] = info + + def _record_loc(self, module, name, node): + try: + if type(node).__name__ == "Token": + tok = node + else: + tok = node.getFirstToken() + if tok is not None: + self.loc[(module, name)] = ( + os.path.basename(self.sm.getFileName(tok.location)), + self.sm.getLineNumber(tok.location)) + except Exception: + pass + + def _scan_module(self, mod): + name = _txt(mod.header.name) + info = dict(name=name, ports=[], nets=[], assigns=[], always=[], insts=[]) + self._record_loc(name, name, mod.header.name) + + # ---- ports (ANSI) ---- + pl = mod.header.ports + if pl is not None: + for p in pl: + if type(p).__name__ == "Token": + continue + try: + d = getattr(p, "declarator", None) + if d is None or d.name is None: + # explicit port like .name(expr) or non-ansi header decl — skip here + continue + pname = _txt(d.name) + h = p.header + dtok = getattr(h, "direction", None) + dtext = (_txt(dtok).capitalize() if dtok is not None else "In") + info["ports"].append((dtext, pname)) + self._record_loc(name, pname, d.name) + except Exception: + continue + + # ---- members ---- + for m in mod.members: + k = m.kind + if k in (sl.SyntaxKind.DataDeclaration, sl.SyntaxKind.NetDeclaration): + for d in m.declarators: + if type(d).__name__ == "Token": + continue + dn = getattr(d, "name", None) + if dn is not None: + n = _txt(dn) + info["nets"].append(n) + self._record_loc(name, n, dn) + elif k == sl.SyntaxKind.ContinuousAssign: + for a in m.assignments: + if type(a).__name__ == "Token": + continue + lhs_ids = _lhs_names(a.left) if getattr(a, "left", None) is not None else [] + rhs_ids = ids_from(a.right) if a.right is not None else [] + if lhs_ids: + info["assigns"].append((lhs_ids[0], rhs_ids)) + self._record_loc(name, lhs_ids[0], m.getFirstToken()) + elif k in (sl.SyntaxKind.AlwaysBlock, sl.SyntaxKind.AlwaysFFBlock, + sl.SyntaxKind.AlwaysCombBlock, sl.SyntaxKind.AlwaysLatchBlock): + self._scan_always(name, m, info) + elif k == sl.SyntaxKind.HierarchyInstantiation: + self._scan_inst(name, m, info) + info["nets"] = sorted(set(info["nets"])) + return info + + def _scan_always(self, mod_name, m, info): + if getattr(m, "statement", None) is None: + return + def walk_stmt(s, depth=0): + if s is None or depth > 40: + return + k = s.kind + if k == sl.SyntaxKind.ExpressionStatement: + e = s.expr + ek = str(e.kind) + if any(a in ek for a in ASSIGN_KINDS): + lt = _syntax_of(e.left) if e.left is not None else None + rt = _syntax_of(e.right) if e.right is not None else None + lhs = _lhs_names(lt) if lt is not None else [] + rhs = ids_from(rt) if rt is not None else [] + if lhs: + info["always"].append((lhs[0], rhs)) + self._record_loc(mod_name, lhs[0], s.getFirstToken()) + return + # descend through timing/if/for/case wrappers + inner = getattr(s, "statement", None) + if inner is not None and hasattr(inner, "kind"): + walk_stmt(inner, depth + 1) + if k == sl.SyntaxKind.ConditionalStatement: + cons = getattr(s, "statement", None) + alt = getattr(s, "elseStatement", None) + for part in (cons, alt): + if part is not None and hasattr(part, "kind"): + walk_stmt(part, depth + 1) + return + if hasattr(s, "__iter__"): + for c in s: + tn = type(c).__name__ + if hasattr(c, "__iter__") and ("Statement" in tn or "Case" in tn + or "Generate" in tn or "Block" in tn): + walk_stmt(c, depth + 1) + walk_stmt(m.statement) + + def _scan_inst(self, mod_name, m, info): + tnode = m.type + itype = _txt(tnode.valueText if hasattr(tnode, "valueText") else tnode) + itype = itype.split()[-1] if itype.split() else "?" + for inst in m.instances: + conns = [] + try: + for pc in inst.connections: + if type(pc).__name__ == "Token": + continue + nm = getattr(pc, "name", None) + if nm is None: + continue + pname = _txt(nm) + expr_ids = ids_from(pc.expr) if getattr(pc, "expr", None) is not None else [] + conns.append((pname, expr_ids)) + except Exception: + pass + iname = _txt(inst.decl.name) if inst.decl is not None and inst.decl.name else "?" + info["insts"].append((itype, iname, conns)) + + # ---------- finalize ---------- + def finalize(self): + self.edges = [] + for mod in self.db.values(): + for (itype, iname, conns) in mod["insts"]: + child = self.db.get(itype) + if not child: + continue + cports = {pn: d for d, pn in child["ports"]} + for (pname, pids) in conns: + cdir = cports.get(pname) + if cdir is None or not pids: + continue + for pid in pids: + self.edges.append((mod["name"], pid, itype, pname, cdir, iname)) + # adjacency for speed + self.in_by_child = {} + for e in self.edges: + if e[4].lower().startswith("in"): + self.in_by_child.setdefault((e[2], e[3]), []).append(e) + self.out_by_child = {} + for e in self.edges: + if e[4].lower().startswith("out"): + self.out_by_child.setdefault((e[2], e[3]), []).append(e) + + # ---------- queries ---------- + def drivers(self, module, sig, max_depth=4, max_nodes=60): + seen, out, rowseen = set(), [], set() + def rec(mod_name, s, depth): + if len(out) >= max_nodes or depth > max_depth or (mod_name, s) in seen: + return + seen.add((mod_name, s)) + m = self.db.get(mod_name) + if not m: + return + for lhs, rhs in m["assigns"]: + if lhs == s: + loc = self.loc.get((mod_name, lhs), ("?", "?")) + out.append(dict(kind="assign", where=f"{mod_name} ({loc[0]}:{loc[1]})", + sig=lhs, rhs=rhs[:8])) + for r in rhs[:3]: + rec(mod_name, r, depth + 1) + for lhs, rhs in m["always"]: + if lhs == s: + loc = self.loc.get((mod_name, lhs), ("?", "?")) + out.append(dict(kind="always", where=f"{mod_name} ({loc[0]}:{loc[1]})", + sig=lhs, rhs=rhs[:8])) + for r in rhs[:3]: + rec(mod_name, r, depth + 1) + pdirs = {pn: d for d, pn in m["ports"]} + if str(pdirs.get(s, "")).lower().startswith("in"): + for (pm, pid, cm, cp, cdir, iname) in self.in_by_child.get((mod_name, s), []): + out.append(dict(kind="port-in", where=f"from {pm}.{iname}", + sig=f".{cp}", rhs=[pid])) + rec(pm, pid, depth + 1) + for (itype, iname, conns) in m["insts"]: + for (pname, pids) in conns: + if s in pids: + child = self.db.get(itype) + if child: + cdirs = {pn: d for d, pn in child["ports"]} + if str(cdirs.get(pname, "")).lower().startswith("out"): + out.append(dict(kind="inst-out", + where=f"{mod_name}.{iname} ({itype}.{pname})", + sig=s, rhs=[f"{itype}.{pname}"])) + rec(itype, pname, depth + 1) + rec(module, sig, 0) + return [r for r in out + if not (key := (r['kind'], r['where'], r['sig'], tuple(r['rhs']))) in rowseen and not rowseen.add(key)] + + def loads(self, module, sig, max_depth=3, max_nodes=60): + """who reads sig (same module) and where does it flow (output port / down into insts)""" + seen, out, rowseen = set(), [], set() + def rec(mod_name, s, depth): + if len(out) >= max_nodes or depth > max_depth or (mod_name, s) in seen: + return + seen.add((mod_name, s)) + m = self.db.get(mod_name) + if not m: + return + for lhs, rhs in m["assigns"] + m["always"]: + if s in rhs: + loc = self.loc.get((mod_name, lhs), ("?", "?")) + out.append(dict(kind="load", where=f"{mod_name} ({loc[0]}:{loc[1]})", + sig=lhs, rhs=[s])) + rec(mod_name, lhs, depth + 1) + pdirs = {pn: d for d, pn in m["ports"]} + if str(pdirs.get(s, "")).lower().startswith("out"): + for (pm, pid, cm, cp, cdir, iname) in self.out_by_child.get((mod_name, s), []): + out.append(dict(kind="port-out", where=f"into {pm} (as {pid}, via {iname})", + sig=f".{cp}", rhs=[pid])) + rec(pm, pid, depth + 1) + for (itype, iname, conns) in m["insts"]: + for (pname, pids) in conns: + if s in pids: + child = self.db.get(itype) + cdirs = {pn: d for d, pn in child["ports"]} if child else {} + kind = cdirs.get(pname, "?") + out.append(dict(kind=f"inst-load-{kind.lower() or 'unk'}", + where=f"{mod_name}.{iname} -> {itype}.{pname}", + sig=s, rhs=[pname])) + rec(module, sig, 0) + return [r for r in out + if not (key := (r['kind'], r['where'], r['sig'], tuple(r['rhs']))) in rowseen and not rowseen.add(key)] + + def module_summary(self, name): + m = self.db.get(name) + if not m: + return None + return dict(name=m["name"], + ports=[f"{d} {p}" for d, p in m["ports"]], + nets=m["nets"][:120], + n_assign=len(m["assigns"]), n_always=len(m["always"]), + insts=[dict(type=t, name=i, ports=[f".{p}({','.join(ids[:3])})" for p, ids in c][:16]) + for t, i, c in m["insts"]]) + + def search(self, pat, limit=20): + rx = re.compile(pat, re.I) + hits = [] + for (mod, sig), (f, l) in self.loc.items(): + if rx.search(sig) and len(hits) < limit: + hits.append(dict(module=mod, signal=sig, file=f, line=l)) + return hits + + def save(self, path): + with open(path, "w", encoding="utf-8") as f: + json.dump(dict(modules=list(self.db), n_edges=len(self.edges)), f) + + @classmethod + def build(cls, roots): + idx = cls() + files = [] + for r in roots: + if os.path.isfile(r): + files.append(r) + else: + files += glob.glob(os.path.join(r, "**", "*.v"), recursive=True) + files += glob.glob(os.path.join(r, "**", "*.sv"), recursive=True) + idx._roots = [os.path.abspath(r) for r in roots] + for f in sorted(set(files)): + try: + idx.add_tree(f) + except Exception: + pass + idx.finalize() + return idx