Commit 71a89423 authored by 方海彤's avatar 方海彤 👶🏻

feat: add cross-file symbol reference context to LLM prompts

Introduces `find_symbol_references` in `ai_review/git.py` to detect deleted symbols within a diff and search for remaining references across the repository using `git grep`. This provides the LLM with factual context to determine if callers have been properly updated, reducing "hallucinated" risks regarding unsynchronized changes.

Key changes:
- Implements heuristic-based symbol extraction for deleted lines (supporting Kotlin, Java, Swift, JS/TS, etc.).
- Filters out references within the files already modified in the current diff.
- Appends a formatted "Symbol Reference Context" block to both standard and deep review prompts.
- Limits analysis to the first 8 symbols and 20 references per symbol to manage prompt size.
parent 7a3f7a55
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
"""Git diff 获取、过滤、统计。""" """Git diff 获取、过滤、统计、跨文件引用检索。"""
import fnmatch import fnmatch
import os import os
import re
import subprocess import subprocess
from ai_review.config import SKIP_PATTERNS from ai_review.config import SKIP_PATTERNS
...@@ -64,3 +65,106 @@ def diff_stats(diff_text): ...@@ -64,3 +65,106 @@ def diff_stats(diff_text):
removed = sum(1 for l in lines if l.startswith("-") and not l.startswith("---")) removed = sum(1 for l in lines if l.startswith("-") and not l.startswith("---"))
files = sum(1 for l in lines if l.startswith("diff --git")) files = sum(1 for l in lines if l.startswith("diff --git"))
return files, added, removed return files, added, removed
# ─── 跨文件符号引用检索 ────────────────────────────────────────
# 让 LLM 不再"假设"调用方有没有同步清理,而是直接看 git grep 结果。
# 常见关键字 / 通用类型,避免被误识别为符号
_REF_KEYWORDS = {
"if", "else", "for", "while", "return", "true", "false", "null", "this",
"new", "var", "val", "let", "const", "fun", "def", "class", "object",
"import", "package", "private", "public", "protected", "internal", "static",
"abstract", "final", "open", "override", "suspend", "inline", "data",
"void", "int", "long", "float", "double", "bool", "boolean", "string",
"self", "init", "deinit", "lazy", "lateinit", "function", "interface",
"enum", "struct", "trait", "extends", "implements", "throws", "throw",
"try", "catch", "finally", "switch", "case", "default", "break", "continue",
"operator", "infix", "companion", "sealed", "annotation", "typealias",
}
# 多语言"删除行可能是定义"的启发式
_DEF_PATTERNS = [
# Kotlin/Java/Swift/Scala: fun/def/class/interface/object/enum/struct NAME
re.compile(r"\b(?:fun|def|class|interface|object|enum|struct|trait)\s+([A-Za-z_][A-Za-z0-9_]*)"),
# Java/C 风格方法: [modifiers] returnType NAME(
re.compile(r"\b(?:public|private|protected|internal|static|final|abstract|override|open)\s+(?:\w+\s+)*([A-Za-z_][A-Za-z0-9_]*)\s*\("),
# JS/TS: function NAME / const NAME = / let NAME =
re.compile(r"\b(?:function|const|let)\s+([A-Za-z_][A-Za-z0-9_]*)"),
# Kotlin 顶层属性: val NAME / var NAME
re.compile(r"\b(?:val|var)\s+([A-Za-z_][A-Za-z0-9_]*)"),
]
_MAX_REFS_PER_SYMBOL = 20
_MAX_SYMBOLS = 8
_MIN_SYMBOL_LEN = 4
def _extract_removed_symbols(diff_text):
"""从 diff 删除行里抽出可能的符号名(保持出现顺序)。"""
seen = set()
ordered = []
for line in diff_text.splitlines():
if not line.startswith("-") or line.startswith("---"):
continue
content = line[1:].strip()
if not content:
continue
for pat in _DEF_PATTERNS:
for m in pat.finditer(content):
name = m.group(1)
if (len(name) < _MIN_SYMBOL_LEN
or name.lower() in _REF_KEYWORDS
or name in seen):
continue
seen.add(name)
ordered.append(name)
return ordered
def find_symbol_references(diff_text, project_root="."):
"""识别 diff 里被删除的符号,在项目其它文件中 grep 引用情况。
返回一段可拼到 prompt 末尾的 Markdown 文本;diff 里没有可分析的删除符号
时返回空串。grep 的结果会过滤掉 diff 自身改动的文件(本就在 diff 里)。
"""
symbols = _extract_removed_symbols(diff_text)
if not symbols:
return ""
symbols = symbols[:_MAX_SYMBOLS]
changed_files = set()
for line in diff_text.splitlines():
if line.startswith("diff --git"):
parts = line.split(" b/")
if len(parts) > 1:
changed_files.add(parts[-1].strip())
blocks = []
for sym in symbols:
try:
result = subprocess.run(
["git", "grep", "-n", "-F", sym],
capture_output=True, text=True, cwd=project_root,
timeout=10,
)
except (subprocess.TimeoutExpired, FileNotFoundError):
continue
external = []
for ln in result.stdout.splitlines():
file_part = ln.split(":", 1)[0]
if file_part in changed_files:
continue
external.append(ln)
if not external:
blocks.append(f"### `{sym}`\n (在其它文件中未找到引用,调用方应已清理)")
else:
shown = external[:_MAX_REFS_PER_SYMBOL]
lines = [f"### `{sym}`"] + [f" {ln}" for ln in shown]
if len(external) > _MAX_REFS_PER_SYMBOL:
lines.append(f" ... 还有 {len(external) - _MAX_REFS_PER_SYMBOL} 处")
blocks.append("\n".join(lines))
return "\n\n".join(blocks)
...@@ -11,7 +11,7 @@ from ai_review.config import ( ...@@ -11,7 +11,7 @@ from ai_review.config import (
STANDARD_TIMEOUT, DEEP_TIMEOUT, STANDARD_TIMEOUT, DEEP_TIMEOUT,
format_duration, format_duration,
) )
from ai_review.git import get_ci_diff, truncate_diff, diff_stats from ai_review.git import get_ci_diff, truncate_diff, diff_stats, find_symbol_references
from ai_review.llm import call_llm, build_standard_prompt, build_deep_prompt from ai_review.llm import call_llm, build_standard_prompt, build_deep_prompt
from ai_review.triage import ( from ai_review.triage import (
build_triage_prompt, parse_triage_decision, build_triage_prompt, parse_triage_decision,
...@@ -20,6 +20,15 @@ from ai_review.triage import ( ...@@ -20,6 +20,15 @@ from ai_review.triage import (
from ai_review.notify import send_feishu from ai_review.notify import send_feishu
_REFS_PROMPT_HEADER = (
"\n\n## 项目内符号引用上下文\n\n"
"以下是 diff 中被删除符号在当前代码库其它位置的引用情况,请优先依据这些事实"
"判断调用方是否同步处理:若某符号"
"「在其它文件中未找到引用」,请直接放行不要列为"
"需要补充上下文的风险;若仍有引用,请明确指出哪些文件未同步并定为高/中风险。\n\n"
)
def main(): def main():
total_start = time.time() total_start = time.time()
...@@ -99,9 +108,15 @@ def main(): ...@@ -99,9 +108,15 @@ def main():
if was_truncated: if was_truncated:
print(f" ⚠️ diff 过长,已截断至 {MAX_DIFF_LINES} 行") print(f" ⚠️ diff 过长,已截断至 {MAX_DIFF_LINES} 行")
refs_block = find_symbol_references(diff)
if refs_block:
print(f" 🔍 已附加 {refs_block.count('### `')} 个符号的项目引用上下文")
print(f" ⏳ 正在调用标准审查 ({TRIAGE_MODEL})...") print(f" ⏳ 正在调用标准审查 ({TRIAGE_MODEL})...")
t_review = time.time() t_review = time.time()
system_msg, user_msg = build_standard_prompt(diff) system_msg, user_msg = build_standard_prompt(diff)
if refs_block:
user_msg += _REFS_PROMPT_HEADER + refs_block
review_result, review_model_used = call_llm( review_result, review_model_used = call_llm(
system_msg, user_msg, system_msg, user_msg,
max_tokens=STANDARD_MAX_TOKENS, max_tokens=STANDARD_MAX_TOKENS,
...@@ -115,9 +130,15 @@ def main(): ...@@ -115,9 +130,15 @@ def main():
if was_truncated: if was_truncated:
print(f" ⚠️ diff 过长,已截断至 {MAX_DIFF_LINES} 行") print(f" ⚠️ diff 过长,已截断至 {MAX_DIFF_LINES} 行")
refs_block = find_symbol_references(diff)
if refs_block:
print(f" 🔍 已附加 {refs_block.count('### `')} 个符号的项目引用上下文")
print(f" ⏳ 正在调用深度审查 ({MODEL}, thinking mode)...") print(f" ⏳ 正在调用深度审查 ({MODEL}, thinking mode)...")
t_review = time.time() t_review = time.time()
system_msg, user_msg = build_deep_prompt(diff) system_msg, user_msg = build_deep_prompt(diff)
if refs_block:
user_msg += _REFS_PROMPT_HEADER + refs_block
review_result, review_model_used = call_llm( review_result, review_model_used = call_llm(
system_msg, user_msg, system_msg, user_msg,
max_tokens=DEEP_MAX_TOKENS, max_tokens=DEEP_MAX_TOKENS,
......
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment