josie / llm-bench

"""Coding-suite instrument — single-turn bug-fix + tool-calling tasks.

Two tiers of difficulty. Each bug-fix answer is compiled and run against
hidden tests with the real toolchain (Go/Rust/Swift/TS/Python). Tool-calling
answers are parse-validated. The suite is pre-validated (buggy provably
fails, reference fix provably passes) by the `validate` command.

Usage (via the bench.py CLI):
  python bench.py coding-suite validate
  python bench.py coding-suite run
  python bench.py coding-suite report

Env (via .env): ENDPOINT, MODEL, API_KEY, MAX_TOKENS, TEMPERATURE.
"""
import json
import os
import re
import sys

from bench.core import client
from bench.core.runners import RUNNERS, extract_code
from bench.core.reporting import jsonl_path, append_jsonl, write_json, coding_suite_summary

HERE = os.path.dirname(os.path.abspath(__file__))

# ---------------------------------------------------------------- task registry
# Tier 1 — straightforward bug-fixes
BUGFIX_T1 = [
    {
        "id": "go-average", "lang": "go", "cat": "bugfix",
        "symptom": "It should return the arithmetic mean as a float, but returns truncated values (e.g. 1 instead of 1.5 for [1,2]).",
        "buggy": "package solution\n\nfunc Average(xs []int) float64 {\n\tsum := 0\n\tfor _, x := range xs {\n\t\tsum += x\n\t}\n\treturn float64(sum / len(xs))\n}\n",
        "fix":   "package solution\n\nfunc Average(xs []int) float64 {\n\tsum := 0\n\tfor _, x := range xs {\n\t\tsum += x\n\t}\n\treturn float64(sum) / float64(len(xs))\n}\n",
        "test":  'package solution\n\nimport "testing"\n\nfunc almost(a, b float64) bool { d := a - b; if d < 0 { d = -d }; return d < 1e-9 }\n\nfunc TestAverage(t *testing.T) {\n\tif !almost(Average([]int{1, 2}), 1.5) { t.Fatalf("[1,2]=%v want 1.5", Average([]int{1, 2})) }\n\tif !almost(Average([]int{2, 2, 2}), 2.0) { t.Fatalf("[2,2,2]=%v want 2", Average([]int{2, 2, 2})) }\n\tif !almost(Average([]int{1, 2, 3, 4}), 2.5) { t.Fatalf("[1,2,3,4]=%v want 2.5", Average([]int{1, 2, 3, 4})) }\n}\n',
    },
    {
        "id": "go-nilmap", "lang": "go", "cat": "bugfix",
        "symptom": "It panics at runtime ('assignment to entry in nil map') instead of grouping the integers by parity.",
        "buggy": 'package solution\n\nfunc GroupByParity(xs []int) map[string][]int {\n\tvar m map[string][]int\n\tfor _, x := range xs {\n\t\tif x%2 == 0 {\n\t\t\tm["even"] = append(m["even"], x)\n\t\t} else {\n\t\t\tm["odd"] = append(m["odd"], x)\n\t\t}\n\t}\n\treturn m\n}\n',
        "fix":   'package solution\n\nfunc GroupByParity(xs []int) map[string][]int {\n\tm := make(map[string][]int)\n\tfor _, x := range xs {\n\t\tif x%2 == 0 {\n\t\t\tm["even"] = append(m["even"], x)\n\t\t} else {\n\t\t\tm["odd"] = append(m["odd"], x)\n\t\t}\n\t}\n\treturn m\n}\n',
        "test":  'package solution\n\nimport (\n\t"reflect"\n\t"testing"\n)\n\nfunc TestGroup(t *testing.T) {\n\tgot := GroupByParity([]int{1, 2, 3, 4})\n\twant := map[string][]int{"odd": {1, 3}, "even": {2, 4}}\n\tif !reflect.DeepEqual(got, want) { t.Fatalf("got %v want %v", got, want) }\n}\n',
    },
    {
        "id": "rust-sumeven", "lang": "rust", "cat": "bugfix",
        "symptom": "sum_even is supposed to sum the EVEN numbers, but it currently sums the odd ones (e.g. returns 4 instead of 6 for [1,2,3,4]).",
        "buggy": "pub fn sum_even(v: &[i32]) -> i32 {\n    v.iter().filter(|&&x| x % 2 == 1).sum()\n}\n",
        "fix":   "pub fn sum_even(v: &[i32]) -> i32 {\n    v.iter().filter(|&&x| x % 2 == 0).sum()\n}\n",
        "test":  "#[cfg(test)]\nmod hidden {\n    use super::*;\n    #[test]\n    fn t() {\n        assert_eq!(sum_even(&[1, 2, 3, 4]), 6);\n        assert_eq!(sum_even(&[2, 4, 6]), 12);\n        assert_eq!(sum_even(&[1, 3, 5]), 0);\n    }\n}\n",
    },
    {
        "id": "rust-safediv", "lang": "rust", "cat": "bugfix",
        "symptom": "safe_div should return none when dividing by zero, but it panics (attempt to divide by zero) instead.",
        "buggy": "pub fn safe_div(a: i32, b: i32) -> Option<i32> {\n    Some(a / b)\n}\n",
        "fix":   "pub fn safe_div(a: i32, b: i32) -> Option<i32> {\n    if b == 0 { None } else { Some(a / b) }\n}\n",
        "test":  "#[cfg(test)]\nmod hidden {\n    use super::*;\n    #[test]\n    fn t() {\n        assert_eq!(safe_div(10, 2), Some(5));\n        assert_eq!(safe_div(7, 0), None);\n        assert_eq!(safe_div(-6, 3), Some(-2));\n    }\n}\n",
    },
    {
        "id": "swift-parseage", "lang": "swift", "cat": "bugfix",
        "symptom": "parseAge should return nil for non-numeric input, but it crashes (force-unwrap of nil) on input like \"abc\".",
        "buggy": "func parseAge(_ s: String) -> Int? {\n    return Int(s)!\n}\n",
        "fix":   "func parseAge(_ s: String) -> Int? {\n    return Int(s)\n}\n",
        "test":  'import Foundation\n\nfunc check(_ cond: Bool, _ msg: String) {\n    if !cond { FileHandle.standardError.write(("FAIL: " + msg + "\\n").data(using: .utf8)!); exit(1) }\n}\n\ncheck(parseAge("34") == 34, "parseAge(34)")\ncheck(parseAge("abc") == nil, "parseAge(abc) should be nil")\ncheck(parseAge("0") == 0, "parseAge(0)")\nprint("OK")\n',
    },
    {
        "id": "swift-median", "lang": "swift", "cat": "bugfix",
        "symptom": "median returns the wrong value because it indexes the array without sorting it first (e.g. median([3,1,2]) returns 1 instead of 2).",
        "buggy": "func median(_ xs: [Double]) -> Double {\n    let n = xs.count\n    if n % 2 == 1 {\n        return xs[n / 2]\n    } else {\n        return (xs[n / 2 - 1] + xs[n / 2]) / 2\n    }\n}\n",
        "fix":   "func median(_ xs: [Double]) -> Double {\n    let s = xs.sorted()\n    let n = s.count\n    if n % 2 == 1 {\n        return s[n / 2]\n    } else {\n        return (s[n / 2 - 1] + s[n / 2]) / 2\n    }\n}\n",
        "test":  'import Foundation\n\nfunc check(_ cond: Bool, _ msg: String) {\n    if !cond { FileHandle.standardError.write(("FAIL: " + msg + "\\n").data(using: .utf8)!); exit(1) }\n}\n\ncheck(median([3, 1, 2]) == 2, "median odd")\ncheck(median([4, 1, 3, 2]) == 2.5, "median even")\ncheck(median([10, 2, 8, 4]) == 6, "median even2")\nprint("OK")\n',
    },
    {
        "id": "ts-sort", "lang": "ts", "cat": "bugfix",
        "symptom": "sortNums sorts numbers lexicographically instead of numerically (e.g. [10,2,1] becomes [1,10,2] instead of [1,2,10]).",
        "buggy": "export function sortNums(a: number[]): number[] {\n  return a.sort();\n}\n",
        "fix":   "export function sortNums(a: number[]): number[] {\n  return [...a].sort((x, y) => x - y);\n}\n",
        "test":  'import { sortNums } from "./solution";\nfunction eq(a: number[], b: number[]) { return a.length === b.length && a.every((v, i) => v === b[i]); }\nlet ok = true;\nif (!eq(sortNums([10, 2, 1]), [1, 2, 10])) { console.error("FAIL sort1"); ok = false; }\nif (!eq(sortNums([3, 30, 1, 2]), [1, 2, 3, 30])) { console.error("FAIL sort2"); ok = false; }\nif (!ok) process.exit(1);\nconsole.log("OK");\n',
    },
    {
        "id": "ts-reduce", "lang": "ts", "cat": "bugfix",
        "symptom": "total throws a TypeError on an empty array because reduce has no initial value; total([]) should be 0.",
        "buggy": "export function total(arr: number[]): number {\n  return arr.reduce((a, b) => a + b);\n}\n",
        "fix":   "export function total(arr: number[]): number {\n  return arr.reduce((a, b) => a + b, 0);\n}\n",
        "test":  'import { total } from "./solution";\nlet ok = true;\nif (total([1, 2, 3]) !== 6) { console.error("FAIL t1"); ok = false; }\nif (total([]) !== 0) { console.error("FAIL empty"); ok = false; }\nif (total([5]) !== 5) { console.error("FAIL t3"); ok = false; }\nif (!ok) process.exit(1);\nconsole.log("OK");\n',
    },
]

# Tier 2 — harder, discriminating bug-fixes (type reasoning, language traps)
BUGFIX_T2 = [
    {
        "id": "go-json-tags", "lang": "go", "cat": "bugfix",
        "symptom": "ToJSON should produce JSON with exactly the keys \"name\", \"emailAddress\", and \"age\", but the name key comes out capitalized and the age field is missing.",
        "buggy": 'package solution\n\nimport "encoding/json"\n\ntype User struct {\n\tName  string\n\tEmail string `json:"emailAddress"`\n\tage   int    `json:"age"`\n}\n\nfunc ToJSON(name, email string, age int) (string, error) {\n\tu := User{Name: name, Email: email, age: age}\n\tb, err := json.Marshal(u)\n\treturn string(b), err\n}\n',
        "fix":   'package solution\n\nimport "encoding/json"\n\ntype User struct {\n\tName  string `json:"name"`\n\tEmail string `json:"emailAddress"`\n\tAge   int    `json:"age"`\n}\n\nfunc ToJSON(name, email string, age int) (string, error) {\n\tu := User{Name: name, Email: email, Age: age}\n\tb, err := json.Marshal(u)\n\treturn string(b), err\n}\n',
        "test":  'package solution\n\nimport (\n\t"encoding/json"\n\t"testing"\n)\n\nfunc TestToJSON(t *testing.T) {\n\ts, err := ToJSON("Ada", "ada@x.com", 36)\n\tif err != nil { t.Fatal(err) }\n\tvar m map[string]interface{}\n\tif err := json.Unmarshal([]byte(s), &m); err != nil { t.Fatal(err) }\n\tif m["name"] != "Ada" { t.Fatalf("name key wrong; got map %v", m) }\n\tif m["emailAddress"] != "ada@x.com" { t.Fatalf("email key wrong; got %v", m) }\n\tif m["age"] != float64(36) { t.Fatalf("age key wrong/missing; got %v", m) }\n}\n',
    },
    {
        "id": "go-generics", "lang": "go", "cat": "bugfix",
        "symptom": "MapSlice should return a new slice with f applied to each element, but the result has extra leading zero values (e.g. [0 0 0 2 4 6] instead of [2 4 6]).",
        "buggy": "package solution\n\nfunc MapSlice[T any, U any](xs []T, f func(T) U) []U {\n\tresult := make([]U, len(xs))\n\tfor _, x := range xs {\n\t\tresult = append(result, f(x))\n\t}\n\treturn result\n}\n",
        "fix":   "package solution\n\nfunc MapSlice[T any, U any](xs []T, f func(T) U) []U {\n\tresult := make([]U, 0, len(xs))\n\tfor _, x := range xs {\n\t\tresult = append(result, f(x))\n\t}\n\treturn result\n}\n",
        "test":  'package solution\n\nimport (\n\t"reflect"\n\t"testing"\n)\n\nfunc TestMapSlice(t *testing.T) {\n\tgot := MapSlice([]int{1, 2, 3}, func(x int) int { return x * 2 })\n\tif !reflect.DeepEqual(got, []int{2, 4, 6}) { t.Fatalf("got %v want [2 4 6]", got) }\n\tgs := MapSlice([]int{1, 2}, func(x int) string { return string(rune(\'a\' + x)) })\n\tif !reflect.DeepEqual(gs, []string{"b", "c"}) { t.Fatalf("got %v want [b c]", gs) }\n}\n',
    },
    {
        "id": "rust-sumall", "lang": "rust", "cat": "bugfix",
        "symptom": "sum_all does not compile because the generic type T is unconstrained. It should sum any slice of numbers (integers or floats) and return the total.",
        "buggy": "pub fn sum_all<T>(items: &[T]) -> T {\n    let mut total = 0;\n    for x in items {\n        total += x;\n    }\n    total\n}\n",
        "fix":   "pub fn sum_all<T: Copy + std::iter::Sum<T>>(items: &[T]) -> T {\n    items.iter().copied().sum()\n}\n",
        "test":  "#[cfg(test)]\nmod hidden {\n    use super::*;\n    #[test]\n    fn t() {\n        assert_eq!(sum_all(&[1, 2, 3]), 6);\n        assert_eq!(sum_all(&[10, 20]), 30);\n        let f: f64 = sum_all(&[1.5, 2.5]);\n        assert!((f - 4.0).abs() < 1e-9);\n    }\n}\n",
    },
    {
        "id": "rust-dedup", "lang": "rust", "cat": "bugfix",
        "symptom": "dedup_sorted should sort the vector and remove consecutive duplicates in place, but it panics with an index-out-of-bounds error.",
        "buggy": "pub fn dedup_sorted(v: &mut Vec<i32>) {\n    v.sort();\n    for i in 1..v.len() {\n        if v[i] == v[i - 1] {\n            v.remove(i);\n        }\n    }\n}\n",
        "fix":   "pub fn dedup_sorted(v: &mut Vec<i32>) {\n    v.sort();\n    v.dedup();\n}\n",
        "test":  "#[cfg(test)]\nmod hidden {\n    use super::*;\n    #[test]\n    fn t() {\n        let mut a = vec![3, 1, 2, 2, 3, 1];\n        dedup_sorted(&mut a);\n        assert_eq!(a, vec![1, 2, 3]);\n        let mut b = vec![5, 5, 5];\n        dedup_sorted(&mut b);\n        assert_eq!(b, vec![5]);\n    }\n}\n",
    },
    {
        "id": "swift-generic-max", "lang": "swift", "cat": "bugfix",
        "symptom": "maxElement does not compile because the generic type T is not constrained. It should return the largest element, or nil for an empty array.",
        "buggy": "func maxElement<T>(_ xs: [T]) -> T? {\n    guard !xs.isEmpty else { return nil }\n    var m = xs[0]\n    for x in xs {\n        if x > m { m = x }\n    }\n    return m\n}\n",
        "fix":   "func maxElement<T: Comparable>(_ xs: [T]) -> T? {\n    guard !xs.isEmpty else { return nil }\n    var m = xs[0]\n    for x in xs {\n        if x > m { m = x }\n    }\n    return m\n}\n",
        "test":  'import Foundation\n\nfunc check(_ cond: Bool, _ msg: String) {\n    if !cond { FileHandle.standardError.write(("FAIL: " + msg + "\\n").data(using: .utf8)!); exit(1) }\n}\n\ncheck(maxElement([3, 1, 2]) == 3, "max ints")\ncheck(maxElement([Int]()) == nil, "max empty")\ncheck(maxElement(["b", "a", "c"]) == "c", "max strings")\nprint("OK")\n',
    },
    {
        "id": "swift-mutating", "lang": "swift", "cat": "bugfix",
        "symptom": "This code does not compile: a struct method modifies a stored property, and the calling code prevents mutation. After fixing, runCounter(5) must return 5.",
        "buggy": "struct Counter {\n    var count = 0\n    func increment() {\n        count += 1\n    }\n}\n\nfunc runCounter(_ times: Int) -> Int {\n    let c = Counter()\n    for _ in 0..<times { c.increment() }\n    return c.count\n}\n",
        "fix":   "struct Counter {\n    var count = 0\n    mutating func increment() {\n        count += 1\n    }\n}\n\nfunc runCounter(_ times: Int) -> Int {\n    var c = Counter()\n    for _ in 0..<times { c.increment() }\n    return c.count\n}\n",
        "test":  'import Foundation\n\nfunc check(_ cond: Bool, _ msg: String) {\n    if !cond { FileHandle.standardError.write(("FAIL: " + msg + "\\n").data(using: .utf8)!); exit(1) }\n}\n\ncheck(runCounter(5) == 5, "runCounter(5)")\ncheck(runCounter(0) == 0, "runCounter(0)")\nprint("OK")\n',
    },
    {
        "id": "ts-matrix-alias", "lang": "ts", "cat": "bugfix",
        "symptom": "makeMatrix should return independent rows, but every row is the SAME array reference (writing to one row changes them all).",
        "buggy": "export function makeMatrix(rows: number, cols: number): number[][] {\n  return new Array(rows).fill(new Array(cols).fill(0));\n}\n",
        "fix":   "export function makeMatrix(rows: number, cols: number): number[][] {\n  return Array.from({ length: rows }, () => new Array(cols).fill(0));\n}\n",
        "test":  'import { makeMatrix } from "./solution";\nconst m = makeMatrix(2, 2);\nm[0][0] = 5;\nif (m[1][0] !== 0) { console.error("FAIL aliasing", m); process.exit(1); }\nif (m.length !== 2 || m[0].length !== 2) { console.error("FAIL shape", m); process.exit(1); }\nconsole.log("OK");\n',
    },
    {
        "id": "ts-closure-var", "lang": "ts", "cat": "bugfix",
        "symptom": "makeAdders should return functions that return 0, 1, and 2, but every function returns 3 (they all share the same loop variable).",
        "buggy": "export function makeAdders(): Array<() => number> {\n  const fns: Array<() => number> = [];\n  for (var i = 0; i < 3; i++) {\n    fns.push(() => i);\n  }\n  return fns;\n}\n",
        "fix":   "export function makeAdders(): Array<() => number> {\n  const fns: Array<() => number> = [];\n  for (let i = 0; i < 3; i++) {\n    fns.push(() => i);\n  }\n  return fns;\n}\n",
        "test":  'import { makeAdders } from "./solution";\nconst vals = makeAdders().map((f) => f());\nfunction eq(a: number[], b: number[]) { return a.length === b.length && a.every((v, i) => v === b[i]); }\nif (!eq(vals, [0, 1, 2])) { console.error("FAIL", vals); process.exit(1); }\nconsole.log("OK");\n',
    },
]

BUGFIX = BUGFIX_T1 + BUGFIX_T2

# ---------------------------------------------------------------- tool-calling
WEATHER_TOOL = [{"type": "function", "function": {
    "name": "get_weather",
    "description": "Get the current weather for a location.",
    "parameters": {"type": "object", "properties": {
        "location": {"type": "string", "description": "City name"},
        "unit": {"type": "string", "enum": ["celsius", "fahrenheit"]}},
        "required": ["location", "unit"]}}}]
CALC_TOOL = [{"type": "function", "function": {
    "name": "calculator",
    "description": "Evaluate an arithmetic expression.",
    "parameters": {"type": "object", "properties": {"expression": {"type": "string"}},
        "required": ["expression"]}}}]
WEB_TOOL = {"type": "function", "function": {
    "name": "web_search", "description": "Search the web.",
    "parameters": {"type": "object", "properties": {"query": {"type": "string"}, "limit": {"type": "integer"}},
        "required": ["query"]}}}
THREE_TOOLS = [WEATHER_TOOL[0], CALC_TOOL[0], WEB_TOOL]
TIMER_TOOL = [{"type": "function", "function": {
    "name": "set_timer", "description": "Set a countdown timer.",
    "parameters": {"type": "object", "properties": {"seconds": {"type": "integer", "description": "duration in seconds"}},
        "required": ["seconds"]}}}]


def _v_native(resp):
    tcs = client.tool_calls(resp)
    if not tcs:
        return False, "no tool_calls returned (content=%r)" % client.content(resp)[:120]
    fn = tcs[0]["function"]
    if fn["name"] != "get_weather":
        return False, "wrong tool: %s" % fn["name"]
    try:
        args = json.loads(fn["arguments"])
    except Exception:
        return False, "args not valid JSON: %r" % fn["arguments"][:120]
    loc_ok = "paris" in str(args.get("location", "")).lower()
    unit_ok = str(args.get("unit", "")).lower() == "celsius"
    return (loc_ok and unit_ok), "args=%s loc_ok=%s unit_ok=%s" % (args, loc_ok, unit_ok)


# The XML tool-call format uses angle brackets around a JSON object.
# We build the regex pattern and system prompt here to avoid the literal
# string being misinterpreted by tooling.
_XML_PATTERN = re.compile(r"[\<\u003c]\s*(\{.*?\})\s*[\>\u003e]", re.DOTALL)

def _v_xml(resp):
    c = client.content(resp)
    m = _XML_PATTERN.search(c)
    if not m:
        return False, "no XML tool-call found: %r" % c[:160]
    try:
        obj = json.loads(m.group(1))
    except Exception:
        return False, "inner JSON did not parse: %r" % m.group(1)[:160]
    args = obj.get("arguments", obj.get("parameters", {}))
    name_ok = obj.get("name") == "web_search"
    q_ok = "gemma" in str(args.get("query", "")).lower()
    limit_ok = str(args.get("limit")) == "5"
    return (name_ok and q_ok and limit_ok), "obj=%s name_ok=%s q_ok=%s limit_ok=%s" % (obj, name_ok, q_ok, limit_ok)


def _v_json(resp):
    c = client.content(resp).strip()
    m = re.search(r"\{.*\}", c, re.DOTALL)
    if not m:
        return False, "no JSON object in output: %r" % c[:160]
    try:
        obj = json.loads(m.group(0))
    except Exception as e:
        return False, "JSON parse error: %s | %r" % (e, m.group(0)[:160])
    name_ok = isinstance(obj.get("name"), str) and "maria" in obj["name"].lower()
    age_ok = obj.get("age") == 34
    sk = obj.get("skills")
    skills_ok = isinstance(sk, list) and {"go", "rust"} <= {str(x).lower() for x in sk}
    return (name_ok and age_ok and skills_ok), "obj=%s name=%s age=%s skills=%s" % (obj, name_ok, age_ok, skills_ok)


def _v_notool(resp):
    tcs = client.tool_calls(resp)
    if tcs:
        return False, "incorrectly called a tool: %s" % tcs[0]["function"]["name"]
    c = client.content(resp).lower()
    return ("paris" in c), "answered in text, mentions Paris=%s" % ("paris" in c)


def _v_multi(resp):
    tcs = client.tool_calls(resp)
    locs = []
    for tc in tcs:
        try:
            a = json.loads(tc["function"]["arguments"])
        except Exception:
            a = {}
        locs.append(str(a.get("location", "")).lower())
    blob = " ".join(locs)
    ok = len(tcs) >= 2 and "paris" in blob and "tokyo" in blob
    return ok, "n_calls=%d locs=%s" % (len(tcs), locs)


def _v_select(resp):
    tcs = client.tool_calls(resp)
    if not tcs:
        return False, "no tool call; content=%r" % client.content(resp)[:140]
    fn = tcs[0]["function"]
    if fn["name"] != "calculator":
        return False, "picked %s" % fn["name"]
    try:
        a = json.loads(fn["arguments"])
    except Exception:
        return False, "bad args %r" % fn["arguments"][:120]
    expr = str(a.get("expression", ""))
    return ("47" in expr and "89" in expr), "expr=%r" % expr


def _v_nested(resp):
    c = client.content(resp).strip()
    m = re.search(r"\{.*\}", c, re.DOTALL)
    if not m:
        return False, "no json: %r" % c[:140]
    try:
        o = json.loads(m.group(0))
    except Exception as e:
        return False, "parse err %s" % e
    svc = o.get("service", {})
    ok = (isinstance(svc, dict) and str(svc.get("name", "")).lower() == "api"
          and svc.get("port") == 8080 and o.get("replicas") == 3
          and "prod" in str(o.get("env", "")).lower())
    return ok, "obj=%s" % o


def _v_derived(resp):
    tcs = client.tool_calls(resp)
    if not tcs:
        return False, "no tool call; content=%r" % client.content(resp)[:140]
    fn = tcs[0]["function"]
    if fn["name"] != "set_timer":
        return False, "picked %s" % fn["name"]
    try:
        a = json.loads(fn["arguments"])
    except Exception:
        return False, "bad args"
    return a.get("seconds") == 150, "args=%s" % a


# Build the XML tool-call system prompt without a literal angle-bracket tag.
# The model is told to emit: <JSON object> on one line.
_XML_SYS = ("You can call tools. To call a tool you MUST output exactly one line "
            "of the form " + chr(60) + "{\"name\": \"<tool>\", \"arguments\": {...}}" + chr(62) +
            " and nothing else — no prose, no markdown. Available tool: web_search(query: string, limit: integer).")


TOOLCALL = [
    {"id": "tc-native", "cat": "toolcall", "tools": WEATHER_TOOL, "validator": _v_native,
     "messages": [{"role": "user", "content": "What is the current weather in Paris? Use celsius."}]},
    {"id": "tc-xml", "cat": "toolcall", "tools": None, "validator": _v_xml,
     "messages": [
        {"role": "system", "content": _XML_SYS},
        {"role": "user", "content": "Search the web for 'gemma benchmarks' and return at most 5 results."}]},
    {"id": "tc-json", "cat": "toolcall", "tools": None, "validator": _v_json,
     "messages": [
        {"role": "system", "content": "Output ONLY a single JSON object and nothing else. No markdown fences, no commentary."},
        {"role": "user", "content": "Extract this person into JSON with keys name (string), age (integer), skills (array of strings): 'Maria is 34 years old and knows Go and Rust.'"}]},
    {"id": "tc-notool", "cat": "toolcall", "tools": CALC_TOOL, "validator": _v_notool,
     "messages": [{"role": "user", "content": "What is the capital of France? Answer in one word."}]},
    {"id": "tc-multi", "cat": "toolcall", "tools": WEATHER_TOOL, "validator": _v_multi,
     "messages": [{"role": "user", "content": "What's the weather in Paris and in Tokyo right now? Use celsius for both."}]},
    {"id": "tc-select", "cat": "toolcall", "tools": THREE_TOOLS, "validator": _v_select,
     "messages": [{"role": "user", "content": "Using the tools available to you, compute 47 * 89 and give me the exact product."}]},
    {"id": "tc-nested-json", "cat": "toolcall", "tools": None, "validator": _v_nested,
     "messages": [
        {"role": "system", "content": "Output ONLY a single JSON object, no markdown, no commentary."},
        {"role": "user", "content": "Produce a deployment config JSON with this shape: {\"service\": {\"name\": string, \"port\": integer}, \"replicas\": integer, \"env\": string}. The service is named 'api', listens on port 8080, runs 3 replicas, in the production environment."}]},
    {"id": "tc-derived", "cat": "toolcall", "tools": TIMER_TOOL, "validator": _v_derived,
     "messages": [{"role": "user", "content": "Set a timer for two and a half minutes."}]},
]


def all_tasks():
    return TOOLCALL + BUGFIX


# ---------------------------------------------------------------- phases
def validate():
    print("=== VALIDATION: proving each bug-fix task is well-formed ===\n")
    allok = True
    for t in BUGFIX:
        runner = RUNNERS[t["lang"]]
        bug_pass, bug_out = runner(t, t["buggy"])
        fix_pass, fix_out = runner(t, t["fix"])
        valid = (not bug_pass) and fix_pass
        allok = allok and valid
        status = "OK " if valid else "BAD"
        print("[%s] %-16s buggy_fails=%-5s fix_passes=%-5s" %
              (status, t["id"], (not bug_pass), fix_pass))
        if not valid:
            if bug_pass:
                print("     !! buggy code unexpectedly PASSED tests")
            if not fix_pass:
                print("     !! reference fix FAILED:\n     " + fix_out.replace("\n", "\n     ")[:800])
    print("\nVALIDATION %s" % ("PASSED — all tasks well-formed" if allok else "FAILED"))
    return allok


def run():
    tasks = all_tasks()
    path = jsonl_path(client.MODEL, prefix="coding-suite")
    open(path, "w").close()
    results = []
    print("=== RUN: %d tasks against %s ===\n" % (len(tasks), client.MODEL), flush=True)
    for i, t in enumerate(tasks, 1):
        print("[%d/%d] %-16s ..." % (i, len(tasks), t["id"]), end=" ", flush=True)
        rec = {"id": t["id"], "cat": t["cat"]}
        try:
            if t["cat"] == "toolcall":
                resp, dt = client.call_model(t["messages"], tools=t.get("tools"))
                passed, detail = t["validator"](resp)
            else:
                sysmsg = ("You are an expert %s developer. Fix the bug in the code. "
                          "Respond with ONLY the corrected code in a single fenced code block. "
                          "Keep the same names and signatures. Do NOT add tests, a main function, "
                          "explanations, or extra declarations." % t["lang"].upper())
                user = ("The following %s code has a bug. %s\n\n```%s\n%s```\n\nReturn the corrected code."
                        % (t["lang"], t["symptom"], t["lang"], t["buggy"]))
                resp, dt = client.call_model(
                    [{"role": "system", "content": sysmsg},
                     {"role": "user", "content": user}])
                code = extract_code(client.content(resp))
                rec["code"] = code
                passed, detail = RUNNERS[t["lang"]](t, code)
            u = client.usage(resp)
            rec.update({
                "passed": bool(passed),
                "detail": detail[:1200],
                "latency_s": round(dt, 1),
                "finish": client.finish_reason(resp),
                "total_tokens": u["prompt_tokens"] + u["completion_tokens"],
                "completion_tokens": u["completion_tokens"],
                "reasoning_tokens": u["reasoning_tokens"],
                "lang": t.get("lang", "-"),
            })
            print("%s  (%.0fs, %s tok)" % ("PASS" if passed else "FAIL", dt, rec.get("total_tokens")), flush=True)
        except Exception as e:
            rec.update({"passed": False, "detail": "HARNESS/REQUEST ERROR: %r" % e, "lang": t.get("lang", "-")})
            print("ERROR: %r" % e, flush=True)
        results.append(rec)
        append_jsonl(path, rec)
    coding_suite_summary(results)
    write_json(results, client.MODEL, prefix="coding-suite")


def report():
    path = jsonl_path(client.MODEL, prefix="coding-suite")
    results = [json.loads(l) for l in open(path)]
    coding_suite_summary(results)




def main():
    mode = sys.argv[1] if len(sys.argv) > 1 else 'validate'
    if mode == 'validate':
        sys.exit(0 if validate() else 1)
    elif mode == 'run':
        run()
    elif mode == 'report':
        report()
    else:
        print(f'usage: python bench.py coding-suite [validate|run|report]')
        sys.exit(1)


if __name__ == "__main__":
    main()