File size: 10,789 Bytes
a071401
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
#!/usr/bin/env python3
"""Evaluate BioinfoMCP conversion outputs against gold MCP source servers."""

from __future__ import annotations

import argparse
import ast
import json
import math
from pathlib import Path
from typing import Any

DEFAULT_BENCHMARK = Path(
    "/225040511/project/Hypo_Bio_OS/experiments/bioinfomcp_benchmark/configs/benchmark_subset_500.json"
)
DEFAULT_OUTPUT = Path(
    "/225040511/project/Hypo_Bio_OS/experiments/bioinfomcp_benchmark/results/conversion_metrics.json"
)


def normalize(text: str) -> str:
    return "".join(ch.lower() for ch in text if ch.isalnum())


def effective_code_lines_from_text(text: str) -> int:
    count = 0
    for line in text.splitlines():
        stripped = line.strip()
        if not stripped or stripped.startswith("#"):
            continue
        count += 1
    return count


def extract_functions_from_code(code: str) -> tuple[bool, bool, list[dict[str, Any]], str | None]:
    try:
        tree = ast.parse(code)
    except SyntaxError as exc:
        return False, False, [], str(exc)

    functions = []
    has_mcp_tool = False
    for node in tree.body:
        if not isinstance(node, ast.FunctionDef):
            continue
        decorators = []
        for dec in node.decorator_list:
            decorators.append(ast.unparse(dec) if hasattr(ast, "unparse") else "")
        decorated = any("mcp.tool" in dec.replace(" ", "") for dec in decorators)
        has_mcp_tool = has_mcp_tool or decorated
        if not decorated:
            continue

        defaults = list(node.args.defaults)
        args = [arg for arg in node.args.args if arg.arg not in {"self", "cls"}]
        first_optional = len(args) - len(defaults)
        params = []
        for idx, arg in enumerate(args):
            params.append(
                {
                    "name": normalize(arg.arg),
                    "annotation": normalize(ast.unparse(arg.annotation)) if arg.annotation is not None else "",
                    "required": idx < first_optional,
                }
            )
        functions.append({"name": normalize(node.name), "params": params})

    return True, has_mcp_tool, functions, None


def load_code(path: Path) -> str:
    return path.read_text(encoding="utf-8", errors="ignore")


def discover_prediction_code(pred_root: Path, server_name: str) -> Path | None:
    candidates = [
        pred_root / f"{server_name}.py",
        pred_root / server_name / "generated.py",
        pred_root / server_name / f"{server_name}_server.py",
        pred_root / f"mcp_{server_name}" / "app" / f"{server_name}_server.py",
    ]
    for path in candidates:
        if path.exists():
            return path
    return None


def discover_usage(pred_root: Path, server_name: str) -> dict[str, Any] | None:
    candidates = [
        pred_root / f"{server_name}.usage.json",
        pred_root / server_name / "usage.json",
        pred_root / server_name / "metadata.json",
    ]
    for path in candidates:
        if path.exists():
            return json.loads(path.read_text(encoding="utf-8"))
    return None


def flatten_params(functions: list[dict[str, Any]]) -> set[tuple[str, str]]:
    items = set()
    for fn in functions:
        for param in fn["params"]:
            items.add((fn["name"], param["name"]))
    return items


def flatten_signatures(functions: list[dict[str, Any]]) -> dict[str, list[tuple[str, str, bool]]]:
    result = {}
    for fn in functions:
        result[fn["name"]] = [(p["name"], p["annotation"], p["required"]) for p in fn["params"]]
    return result


def f1(pred: set[Any], gold: set[Any]) -> tuple[float, float, float]:
    if not pred and not gold:
        return 1.0, 1.0, 1.0
    if not pred:
        return 0.0, 0.0, 0.0
    if not gold:
        return 0.0, 0.0, 0.0
    hit = len(pred & gold)
    precision = hit / len(pred)
    recall = hit / len(gold)
    if precision + recall == 0:
        return precision, recall, 0.0
    return precision, recall, 2 * precision * recall / (precision + recall)


def usage_tokens(usage: dict[str, Any] | None) -> tuple[int | None, int | None, int | None]:
    if not usage:
        return None, None, None
    prompt = usage.get("prompt_tokens")
    completion = usage.get("completion_tokens")
    total = usage.get("total_tokens")
    nested = usage.get("usage")
    if isinstance(nested, dict):
        prompt = prompt if prompt is not None else nested.get("prompt_tokens")
        completion = completion if completion is not None else nested.get("completion_tokens")
        total = total if total is not None else nested.get("total_tokens")
    if total is None and isinstance(prompt, int) and isinstance(completion, int):
        total = prompt + completion
    return prompt, completion, total


def main() -> int:
    parser = argparse.ArgumentParser()
    parser.add_argument("--benchmark", type=Path, default=DEFAULT_BENCHMARK)
    parser.add_argument("--pred-root", type=Path, required=True)
    parser.add_argument("--out", type=Path, default=DEFAULT_OUTPUT)
    args = parser.parse_args()

    benchmark = json.loads(args.benchmark.read_text(encoding="utf-8"))
    results = []

    for item in benchmark["items"]:
        server_name = item["server_name"]
        gold_code = load_code(Path(item["gold_source_path"]))
        pred_path = discover_prediction_code(args.pred_root, server_name)
        usage = discover_usage(args.pred_root, server_name)

        gold_ok, _, gold_functions, gold_err = extract_functions_from_code(gold_code)
        pred_code = load_code(pred_path) if pred_path else ""
        pred_ok, has_mcp_tool, pred_functions, pred_err = extract_functions_from_code(pred_code) if pred_path else (
            False,
            False,
            [],
            "missing prediction file",
        )

        gold_tools = {fn["name"] for fn in gold_functions}
        pred_tools = {fn["name"] for fn in pred_functions}
        tool_precision, tool_recall, tool_f1 = f1(pred_tools, gold_tools)

        gold_params = flatten_params(gold_functions)
        pred_params = flatten_params(pred_functions)
        param_precision, param_recall, param_f1 = f1(pred_params, gold_params)

        gold_signatures = flatten_signatures(gold_functions)
        pred_signatures = flatten_signatures(pred_functions)
        exact_signature_matches = sum(
            1 for name, signature in gold_signatures.items() if pred_signatures.get(name) == signature
        )
        signature_exact_rate = exact_signature_matches / len(gold_signatures) if gold_signatures else 0.0

        structural_pass = 1.0 if pred_ok and has_mcp_tool else 0.0
        conversion_accuracy = (
            0.25 * structural_pass
            + 0.35 * tool_f1
            + 0.25 * param_f1
            + 0.15 * signature_exact_rate
        )

        prompt_tokens, completion_tokens, total_tokens = usage_tokens(usage)
        code_lines = effective_code_lines_from_text(pred_code) if pred_code else None
        acc_per_1k_tokens = None
        acc_per_100_loc = None
        if isinstance(total_tokens, int) and total_tokens > 0:
            acc_per_1k_tokens = conversion_accuracy * 1000 / total_tokens
        if isinstance(code_lines, int) and code_lines > 0:
            acc_per_100_loc = conversion_accuracy * 100 / code_lines

        results.append(
            {
                "server_name": server_name,
                "category": item["category"],
                "complexity": item["complexity"],
                "gold_tool_count": len(gold_functions),
                "prediction_file": str(pred_path) if pred_path else None,
                "syntax_valid": pred_ok,
                "has_mcp_tool": has_mcp_tool,
                "tool_precision": tool_precision,
                "tool_recall": tool_recall,
                "tool_f1": tool_f1,
                "param_precision": param_precision,
                "param_recall": param_recall,
                "param_f1": param_f1,
                "signature_exact_rate": signature_exact_rate,
                "conversion_accuracy": conversion_accuracy,
                "code_lines": code_lines,
                "prompt_tokens": prompt_tokens,
                "completion_tokens": completion_tokens,
                "total_tokens": total_tokens,
                "accuracy_per_1k_tokens": acc_per_1k_tokens,
                "accuracy_per_100_loc": acc_per_100_loc,
                "gold_error": gold_err,
                "prediction_error": pred_err,
            }
        )

    valid = [row for row in results if row["prediction_file"]]
    def mean(key: str) -> float | None:
        vals = [row[key] for row in valid if isinstance(row.get(key), (int, float))]
        return sum(vals) / len(vals) if vals else None

    acc_vals = [row["conversion_accuracy"] for row in valid if isinstance(row["conversion_accuracy"], (int, float))]
    token_vals = [row["total_tokens"] for row in valid if isinstance(row["total_tokens"], int) and row["total_tokens"] > 0]
    loc_vals = [row["code_lines"] for row in valid if isinstance(row["code_lines"], int) and row["code_lines"] > 0]
    median_tokens = sorted(token_vals)[len(token_vals) // 2] if token_vals else None
    median_loc = sorted(loc_vals)[len(loc_vals) // 2] if loc_vals else None

    for row in valid:
        if median_tokens and median_loc and row["total_tokens"] and row["code_lines"]:
            token_norm = row["total_tokens"] / median_tokens
            loc_norm = row["code_lines"] / median_loc
            row["token_line_efficiency"] = row["conversion_accuracy"] / (0.6 * token_norm + 0.4 * loc_norm)
        else:
            row["token_line_efficiency"] = None

    payload = {
        "benchmark": str(args.benchmark),
        "pred_root": str(args.pred_root),
        "aggregate": {
            "expected_server_count": len(benchmark["items"]),
            "predicted_server_count": len(valid),
            "mean_conversion_accuracy": mean("conversion_accuracy"),
            "mean_tool_f1": mean("tool_f1"),
            "mean_param_f1": mean("param_f1"),
            "mean_signature_exact_rate": mean("signature_exact_rate"),
            "mean_accuracy_per_1k_tokens": mean("accuracy_per_1k_tokens"),
            "mean_accuracy_per_100_loc": mean("accuracy_per_100_loc"),
            "mean_token_line_efficiency": mean("token_line_efficiency"),
            "median_total_tokens": median_tokens,
            "median_code_lines": median_loc,
        },
        "results": results,
    }

    args.out.parent.mkdir(parents=True, exist_ok=True)
    args.out.write_text(json.dumps(payload, indent=2, ensure_ascii=False), encoding="utf-8")
    print(json.dumps(payload["aggregate"], indent=2, ensure_ascii=False))
    return 0


if __name__ == "__main__":
    raise SystemExit(main())