diff --git a/tests/e2e/core/providers/__init__.py b/tests/e2e/core/providers/__init__.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/tests/e2e/core/providers/_native_helpers.py b/tests/e2e/core/providers/_native_helpers.py new file mode 100644 index 000000000000..c6fa61ccad04 --- /dev/null +++ b/tests/e2e/core/providers/_native_helpers.py @@ -0,0 +1,195 @@ +"""Shared harness for the native-dialect provider wire suites (``test_native_*.py``). + +Every scenario drives the REAL ``hermes`` CLI (``python -m hermes_cli.main chat -q ... -Q``) as a +subprocess with a hermetic fake HOME / HERMES_HOME, a config.yaml that selects a native provider, +and the provider's endpoint redirected to a loopback fake from ``tests/fakes/providers/``. Only the +vendor boundary is faked; runtime resolution, the adapter, the agent loop, tools and SQLite are real. + +What a test asserts: the NEXT wire request Hermes sends (captured by the fake), the persisted +``state.db`` rows, and the CLI's user-visible output — never Hermes source text. +""" + +from __future__ import annotations + +import json +import os +import sqlite3 +import subprocess +import sys +import time +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Callable + +import yaml + +REPO_ROOT = Path(__file__).resolve().parents[4] +TURN_TIMEOUT = 180.0 + +_SECRET_ENV_SUFFIXES = ("_API_KEY", "_TOKEN", "_SECRET", "_ACCESS_KEY", "_SESSION_TOKEN") +_PASSTHROUGH_ENV = frozenset({ + "PATH", "LANG", "LANGUAGE", "USER", "LOGNAME", "SHELL", "TMPDIR", "TZ", + "SYSTEMROOT", "SystemRoot", "COMSPEC", "PATHEXT", "WINDIR", "TEMP", "TMP", +}) + + +@dataclass +class NativeHome: + """One hermetic fake HOME with ``HOME/.hermes`` as HERMES_HOME and a project dir as cwd.""" + + root: Path + + @property + def home(self) -> Path: + return self.root / "home" + + @property + def hermes_home(self) -> Path: + return self.home / ".hermes" + + @property + def project(self) -> Path: + return self.root / "project" + + @property + def db_path(self) -> Path: + return self.hermes_home / "state.db" + + def env(self, extra: dict[str, str] | None = None) -> dict[str, str]: + """Allowlisted env: no inherited credentials, HERMES_* or TERMINAL_* can reroute the child.""" + env = { + k: v for k, v in os.environ.items() + if (k in _PASSTHROUGH_ENV or k.startswith("LC_")) and not k.endswith(_SECRET_ENV_SUFFIXES) + } + env.update({ + "HOME": str(self.home), + "HERMES_HOME": str(self.hermes_home), + "PYTHONPATH": str(REPO_ROOT), + "PYTHONUNBUFFERED": "1", + "NO_COLOR": "1", + "TERM": "dumb", + # The child's ~/.hermes/state.db IS the tmp home's db; under a pytest ancestor the + # live-DB guard would refuse it. Documented child escape hatch; path is tmp by construction. + "HERMES_STATE_DB_GUARD_BYPASS": "1", + # Never reach real AWS/GCP metadata endpoints or shared config from a fake home. + "AWS_EC2_METADATA_DISABLED": "true", + "AWS_CONFIG_FILE": str(self.home / ".aws" / "config"), + "AWS_SHARED_CREDENTIALS_FILE": str(self.home / ".aws" / "credentials"), + "NO_GCE_CHECK": "True", + }) + env.update(extra or {}) + return env + + +def make_home(root: Path, model: dict[str, Any], *, env_file: dict[str, str] | None = None, + extra_config: dict[str, Any] | None = None) -> NativeHome: + """Write config.yaml (``model`` block + offline defaults + ``extra_config``) and ``.env``.""" + nh = NativeHome(root) + for d in (nh.hermes_home, nh.project): + d.mkdir(parents=True, exist_ok=True) + cfg: dict[str, Any] = { + "model": model, + "agent": {"api_max_retries": 2}, + "updates": {"check": False}, + "auxiliary": {"title_generation": {"enabled": False}}, + "memory": {"memory_enabled": False, "user_profile_enabled": False}, + } + for key, value in (extra_config or {}).items(): + if isinstance(value, dict) and isinstance(cfg.get(key), dict): + cfg[key].update(value) + else: + cfg[key] = value + (nh.hermes_home / "config.yaml").write_text(yaml.safe_dump(cfg, sort_keys=False), encoding="utf-8") + lines = [f"{k}={v}" for k, v in (env_file or {}).items()] + (nh.hermes_home / ".env").write_text("\n".join(lines) + ("\n" if lines else ""), encoding="utf-8") + return nh + + +@dataclass +class ChatResult: + returncode: int + stdout: str + stderr: str + wall_s: float + + def describe(self) -> str: + return (f"exit={self.returncode} wall={self.wall_s:.1f}s\n--- stdout ---\n{self.stdout[-2500:]}" + f"\n--- stderr ---\n{self.stderr[-4000:]}") + + +def run_chat(nh: NativeHome, prompt: str, *, resume: str | None = None, env: dict[str, str] | None = None, + args: tuple[str, ...] = (), timeout: float = TURN_TIMEOUT) -> ChatResult: + """One real ``hermes chat -q`` turn (optionally ``--resume ``) against the fake provider.""" + argv = [sys.executable, "-m", "hermes_cli.main", "chat", "-q", prompt, "-Q", *args] + if resume: + argv += ["--resume", resume] + started = time.monotonic() + proc = subprocess.run(argv, cwd=nh.project, env=nh.env(env), capture_output=True, text=True, + timeout=timeout, stdin=subprocess.DEVNULL) + return ChatResult(proc.returncode, proc.stdout, proc.stderr, time.monotonic() - started) + + +def _connect(nh: NativeHome) -> sqlite3.Connection: + conn = sqlite3.connect(f"file:{nh.db_path}?mode=ro", uri=True, timeout=10) + conn.row_factory = sqlite3.Row + return conn + + +def session_ids(nh: NativeHome) -> list[str]: + """All session ids, oldest first.""" + if not nh.db_path.exists(): + return [] + with _connect(nh) as conn: + return [r["id"] for r in conn.execute("SELECT id FROM sessions ORDER BY started_at, rowid")] + + +def latest_session(nh: NativeHome) -> str: + ids = session_ids(nh) + assert ids, f"no session persisted in {nh.db_path}" + return ids[-1] + + +def messages(nh: NativeHome, session_id: str | None = None, *, active_only: bool = True) -> list[dict[str, Any]]: + """Persisted message rows (dicts) for ``session_id`` (default: every session), in insertion order.""" + if not nh.db_path.exists(): + return [] + where, params = [], [] + if session_id: + where.append("session_id = ?") + params.append(session_id) + if active_only: + where.append("active = 1") + sql = "SELECT * FROM messages" + (f" WHERE {' AND '.join(where)}" if where else "") + " ORDER BY id" + with _connect(nh) as conn: + return [dict(r) for r in conn.execute(sql, params)] + + +def tool_calls_of(row: dict[str, Any]) -> list[dict[str, Any]]: + raw = row.get("tool_calls") + return json.loads(raw) if raw else [] + + +class KnownSymptom(AssertionError): + """Raised ONLY at a tracked bug's exact symptom. + + Every KNOWN strict xfail uses ``raises=KnownSymptom`` so a harness failure (process death, timeout, + precondition assert, fixture teardown error) fails for real instead of counting as the known bug. + """ + + +def wait_until(predicate: Callable[[], Any], timeout: float, what: str, interval: float = 0.05, + error: type[AssertionError] = AssertionError) -> Any: + """Poll ``predicate`` until truthy or raise ``error`` naming ``what`` (no bare sleeps as synchronization).""" + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + value = predicate() + if value: + return value + time.sleep(interval) + raise error(f"timed out after {timeout}s waiting for {what}") + + +def assert_no_duplicate_assistant_text(rows: list[dict[str, Any]], needle: str) -> None: + """A retried/dropped stream must never persist the same assistant content twice.""" + hits = [r["id"] for r in rows if r["role"] == "assistant" and needle in (r.get("content") or "")] + assert len(hits) <= 1, f"assistant text {needle!r} persisted {len(hits)}x (rows {hits})" diff --git a/tests/e2e/core/providers/test_native_bedrock_converse_faults.py b/tests/e2e/core/providers/test_native_bedrock_converse_faults.py new file mode 100644 index 000000000000..c822abe01175 --- /dev/null +++ b/tests/e2e/core/providers/test_native_bedrock_converse_faults.py @@ -0,0 +1,173 @@ +"""Bedrock Converse / ConverseStream fault semantics through the real ``hermes chat -q`` CLI. + +Each scenario owns a loopback Bedrock fake (``tests/fakes/providers/bedrock_converse.py``) that the real +boto3 client reaches via ``AWS_ENDPOINT_URL_BEDROCK_RUNTIME``; faults are the service's documented ones +in the AWS JSON error shape (``x-amzn-ErrorType`` + ``{"message"}``) or ``:message-type exception`` +event-stream frames, plus connection drops. Assertions: requests counted at the fake, the retried +request's body, ``state.db`` rows (no duplicated assistant content) and what the CLI prints. +""" + +from __future__ import annotations + +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any + +import pytest + +pytest.importorskip("botocore") + +from tests.e2e.core.providers._native_helpers import ( # noqa: E402 + ChatResult, KnownSymptom, NativeHome, assert_no_duplicate_assistant_text, make_home, messages, run_chat, +) +from tests.fakes.providers.bedrock_converse import ( # noqa: E402 + ACCESS_KEY, REGION, SECRET_KEY, Drop, FakeBedrock, HttpError, Reasoning, Reply, StreamException, Text, + Turn, seq, +) + +MODEL = "deepseek.v3-v1:0" +FINAL = "Recovered answer ZEBRA-7731 is complete and whole." +REASONING = "Reasoning for the recovered answer, streamed in several deltas." +VALIDATION_MARK = "FAKE-VALIDATION-9921" +IAM_DENIAL = ("User: arn:aws:iam::123456789012:user/e2e is not authorized to perform: " + "bedrock:InvokeModelWithResponseStream on resource: arn:aws:bedrock:us-east-1::foundation-model/" + + MODEL) + +KNOWN = { + "validation_retried": "#121294 a Bedrock 400 ValidationException is retried and reported as 'temporarily unavailable'", + "eof_before_message_stop": "#109988 a ConverseStream that ends before messageStop is accepted as the answer", +} + + +CHUNK = 7 +# Event index 3 text deltas into the answer: messageStart, reasoning deltas, signature, contentBlockStop. +CUT_MID_TEXT = 1 + -(-len(REASONING) // CHUNK) + 2 + 3 + + +def _answer() -> Turn: + return Turn([Reasoning(REASONING), Text(FINAL, chunk=CHUNK)]) + + +@dataclass +class Scenario: + replies: tuple[Reply, ...] + env: dict[str, str] = field(default_factory=dict) + by_op: dict[str, tuple[Reply, ...]] = field(default_factory=dict) + + +# Each fault precedes a good answer: a retryable fault must be retried into it, a terminal one must not. +SCENARIOS: dict[str, Scenario] = { + # botocore's own retries are off (documented AWS_MAX_ATTEMPTS) so the 429 reaches Hermes' loop. + "throttle_http": Scenario((HttpError("ThrottlingException", "Too many requests, please wait before trying again."), + _answer()), env={"AWS_MAX_ATTEMPTS": "1"}), + "unavailable_http": Scenario((HttpError("ServiceUnavailableException", "Bedrock is unable to process your request."), + _answer()), env={"AWS_MAX_ATTEMPTS": "1"}), + "throttle_in_stream": Scenario((StreamException(_answer(), "throttlingException", + "Too many tokens, please wait before trying again.", after=CUT_MID_TEXT), + _answer())), + "drop_mid_stream": Scenario((Drop(_answer(), after=CUT_MID_TEXT), _answer())), + "eof_before_message_stop": Scenario((Drop(_answer(), after=CUT_MID_TEXT, clean=True), _answer())), + "validation": Scenario((HttpError("ValidationException", f"The model returned the following errors: " + f"malformed input request: {VALIDATION_MARK}"), _answer())), + "stream_denied_falls_back": Scenario((), by_op={ + "ConverseStream": (HttpError("AccessDeniedException", IAM_DENIAL),), "Converse": (_answer(),)}), +} + + +def _responder(sc: Scenario): + if not sc.by_op: + return seq(*sc.replies) + per_op = {op: seq(*replies) for op, replies in sc.by_op.items()} + return lambda rec: per_op[rec["op"]](rec) + + +def _run(name: str, root: Path) -> dict[str, Any]: + sc = SCENARIOS[name] + fake = FakeBedrock(_responder(sc)) + with fake: + nh = make_home(root, {"provider": "bedrock", "default": MODEL, "context_length": 64000}, + env_file={"AWS_ACCESS_KEY_ID": ACCESS_KEY, "AWS_SECRET_ACCESS_KEY": SECRET_KEY, + "AWS_REGION": REGION}) + result = run_chat(nh, "Give me the recovered answer.", + env={"AWS_ENDPOINT_URL_BEDROCK_RUNTIME": fake.endpoint, **sc.env}) + return {"nh": nh, "result": result, "requests": fake.snapshot()} + + +@pytest.fixture(scope="module") +def runs(tmp_path_factory: pytest.TempPathFactory) -> dict[str, dict[str, Any]]: + base = tmp_path_factory.mktemp("bedrock_faults") + with ThreadPoolExecutor(len(SCENARIOS)) as pool: + futures = {name: pool.submit(_run, name, base / name) for name in SCENARIOS} + return {name: fut.result() for name, fut in futures.items()} + + +def _assistant_rows(nh: NativeHome) -> list[dict[str, Any]]: + return [r for r in messages(nh) if r["role"] == "assistant"] + + +def _recovered(run: dict[str, Any], expected_requests: int) -> None: + """Exit 0, the full answer printed once, one assistant row, the retry re-sent the same history.""" + result: ChatResult = run["result"] + assert result.returncode == 0, result.describe() + requests = run["requests"] + assert not [r["rejected"] for r in requests if r.get("rejected")], requests + assert len(requests) == expected_requests, f"fake saw {len(requests)} requests\n{result.describe()}" + assert all(r["body"]["messages"] == requests[0]["body"]["messages"] for r in requests), \ + "the retry did not re-send the original history (partial output leaked into the request?)" + assert result.stdout.count(FINAL) == 1, result.describe() + rows = _assistant_rows(run["nh"]) + assert_no_duplicate_assistant_text(rows, FINAL[:12]) + assert [r["content"] for r in rows] == [FINAL], rows + + +@pytest.mark.parametrize("name", ["throttle_http", "unavailable_http", "throttle_in_stream"]) +def test_retryable_fault_is_retried_once_into_the_answer(name: str, runs: dict[str, Any]) -> None: + _recovered(runs[name], expected_requests=2) + + +def test_mid_stream_drop_retries_without_duplicated_persisted_content(runs: dict[str, Any]) -> None: + run = runs["drop_mid_stream"] + _recovered(run, expected_requests=2) + assert run["requests"][0]["reply"] == "Drop", run["requests"][0] + + +@pytest.mark.xfail(strict=True, raises=KnownSymptom, reason=KNOWN["eof_before_message_stop"]) +def test_stream_ending_before_message_stop_is_not_accepted(runs: dict[str, Any]) -> None: + run = runs["eof_before_message_stop"] + assert run["result"].returncode == 0, run["result"].describe() + assert run["requests"] and run["requests"][0]["reply"] == "Drop", run["requests"] + rows = [r["content"] for r in _assistant_rows(run["nh"])] + sent = len(run["requests"]) + # Symptom: the truncated first stream is persisted as the final answer and never retried. + if sent == 1 and rows and FINAL.startswith(rows[-1]) and rows[-1] != FINAL: + raise KnownSymptom(f"{KNOWN['eof_before_message_stop']}: requests={sent} rows={rows}") + assert (sent, rows) == (2, [FINAL]), f"requests={sent} rows={rows}" + + +@pytest.mark.xfail(strict=True, raises=KnownSymptom, reason=KNOWN["validation_retried"]) +def test_validation_exception_is_surfaced_once_without_retry(runs: dict[str, Any]) -> None: + run = runs["validation"] + result: ChatResult = run["result"] + sent = len(run["requests"]) + assert sent >= 1 and run["requests"][0]["reply"] == "HttpError", run["requests"] + # Symptom: the 400 is retried (the scripted success behind it is reached). + if sent > 1: + raise KnownSymptom(f"{KNOWN['validation_retried']}: fake saw {sent} requests") + shown = result.stdout + result.stderr + assert result.returncode != 0 and FINAL not in shown, result.describe() + assert VALIDATION_MARK in shown and "ValidationException" in shown, result.describe() + assert "temporarily unavailable" not in shown, result.describe() + assert [r["content"] for r in _assistant_rows(run["nh"])] != [FINAL] + + +def test_streaming_iam_denial_falls_back_to_converse(runs: dict[str, Any]) -> None: + run = runs["stream_denied_falls_back"] + result: ChatResult = run["result"] + assert result.returncode == 0, result.describe() + requests = run["requests"] + assert [r["op"] for r in requests] == ["ConverseStream", "Converse"], [r["op"] for r in requests] + assert requests[1]["body"]["messages"] == requests[0]["body"]["messages"] + assert FINAL in result.stdout, result.describe() + rows = _assistant_rows(run["nh"]) + assert [(r["content"], r["reasoning_content"]) for r in rows] == [(FINAL, REASONING)], rows diff --git a/tests/e2e/core/providers/test_native_bedrock_converse_turns.py b/tests/e2e/core/providers/test_native_bedrock_converse_turns.py new file mode 100644 index 000000000000..c95de23f6fd7 --- /dev/null +++ b/tests/e2e/core/providers/test_native_bedrock_converse_turns.py @@ -0,0 +1,365 @@ +"""Bedrock Converse / ConverseStream wire conformance: tools, signed reasoning, resume, compaction. + +Real ``hermes chat -q`` subprocesses with ``model.provider: bedrock`` talk to the real boto3 +``bedrock-runtime`` client, redirected by botocore's documented ``AWS_ENDPOINT_URL_BEDROCK_RUNTIME`` +override to the loopback fake in ``tests/fakes/providers/bedrock_converse.py``. The fake verifies the +SigV4 signature, validates each body against the botocore service model plus Converse's conversation +rules (tool pairing, signed reasoning replay, final-assistant thinking), and streams real +``application/vnd.amazon.eventstream`` frames. Independent scenarios run concurrently in one module +fixture; each test asserts one property of the wire requests, ``state.db`` rows or CLI output. +""" + +from __future__ import annotations + +import json +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from typing import Any, Callable + +import pytest + +pytest.importorskip("botocore") + +from tests.e2e.core.providers._native_helpers import ( # noqa: E402 + ChatResult, KnownSymptom, NativeHome, latest_session, make_home, messages, run_chat, session_ids, tool_calls_of, +) +from tests.fakes.providers.bedrock_converse import ( # noqa: E402 + ACCESS_KEY, REGION, SECRET_KEY, FakeBedrock, Reasoning, Text, ToolUse, Turn, seq, +) + +MODEL = "deepseek.v3-v1:0" + +KNOWN = { + "reasoning_shredded": "#98468 streamed reasoning is persisted with '\\n\\n' between every delta", + "resume_drops_reasoning": "#121293 --resume replays Bedrock assistant turns without their signed reasoningContent", +} + +SEED_TEXT = "codeword PELICAN-5501" +SEED2_TEXT = "second codeword OSPREY-7712" +R_A1 = "The user wants both the seed file and an echo; run them in parallel first." +R_A2 = "Both results are back; now read the second seed before answering." +R_A3 = "I have both codewords and the echo, so I can answer now." +FINAL_A = "Answer: PELICAN-5501 / OSPREY-7712 / MARK-42." +R_B1 = "Resume scenario: read the seed file to learn the codeword before replying." +R_B2 = "The seed file holds the codeword, report it back verbatim to the user." +FINAL_B1 = "The codeword is PELICAN-5501." +FINAL_B2 = "Still PELICAN-5501 after the resume." +SUMMARY_TOKEN = "SUMMARY-OK-BEDROCK" +FINAL_C = "Compaction scenario finished: DONE-C." +COMPACTION_FILES = 8 + + +def _home(root: Path, fake: FakeBedrock, extra_config: dict[str, Any] | None = None) -> NativeHome: + """Bedrock home: fake static creds in the profile .env (the chain Hermes loads), endpoint via env.""" + nh = make_home(root, {"provider": "bedrock", "default": MODEL, "context_length": 64000}, + env_file={"AWS_ACCESS_KEY_ID": ACCESS_KEY, "AWS_SECRET_ACCESS_KEY": SECRET_KEY, + "AWS_REGION": REGION}, + extra_config=extra_config) + (nh.project / "seed.txt").write_text(SEED_TEXT + "\n", encoding="utf-8") + (nh.project / "seed2.txt").write_text(SEED2_TEXT + "\n", encoding="utf-8") + return nh + + +def _endpoint_env(fake: FakeBedrock) -> dict[str, str]: + return {"AWS_ENDPOINT_URL_BEDROCK_RUNTIME": fake.endpoint} + + +# -------------------------------------------------------------------------------------------------- +# Scenarios (each owns its fake + home; run concurrently) +# -------------------------------------------------------------------------------------------------- + + +def _scenario_tools(root: Path) -> dict[str, Any]: + """Parallel toolUse (read_file + terminal) -> toolResults, a second tool round, then the answer.""" + def first(_rec: dict[str, Any]) -> Turn: + return Turn([Reasoning(R_A1), ToolUse("read_file", {"path": str(root / "project" / "seed.txt")}), + ToolUse("terminal", {"command": "seq 42 42 | sed s/^/MARK-/"})]) + + fake = FakeBedrock(seq(first, lambda _r: Turn([Reasoning(R_A2), ToolUse( + "read_file", {"path": str(root / "project" / "seed2.txt")})]), Turn([Reasoning(R_A3), Text(FINAL_A)]))) + with fake: + nh = _home(root, fake) + result = run_chat(nh, "Read seed.txt, echo a marker, then read seed2.txt and report.", + env=_endpoint_env(fake)) + return {"fake": fake, "nh": nh, "result": result, "requests": fake.snapshot()} + + +def _scenario_resume(root: Path) -> dict[str, Any]: + """Turn 1 (reasoning + toolUse, then reasoning + text) in one process; turn 2 via --resume in another.""" + fake = FakeBedrock(seq( + lambda _r: Turn([Reasoning(R_B1), ToolUse("read_file", {"path": str(root / "project" / "seed.txt")})]), + Turn([Reasoning(R_B2), Text(FINAL_B1)]), Turn([Text(FINAL_B2)]))) + with fake: + nh = _home(root, fake) + first = run_chat(nh, "What codeword is in seed.txt?", env=_endpoint_env(fake)) + sid = latest_session(nh) if first.returncode == 0 else "" + second = run_chat(nh, "Say it again.", resume=sid, env=_endpoint_env(fake)) if sid else first + return {"fake": fake, "nh": nh, "first": first, "second": second, "sid": sid, "requests": fake.snapshot()} + + +def _compaction_responder(root: Path) -> Callable[[dict[str, Any]], Turn]: + """Main turns (with toolConfig) read files and report huge input usage; aux calls are summaries.""" + main_calls: list[int] = [] + + def respond(rec: dict[str, Any]) -> Turn: + if "toolConfig" not in rec["body"]: + return Turn([Text(f"## Goal\nRead the files ({SUMMARY_TOKEN}).\n## Progress\n- files read so far\n")], + input_tokens=600) + main_calls.append(1) + n = len(main_calls) + # Small usage first so the transcript is long enough to compact, then far past the threshold. + usage = 40_000 if n >= 6 else 3_000 + if n <= COMPACTION_FILES: + return Turn([Reasoning(f"Step {n}: read f{n}.txt next and keep going through the list."), + ToolUse("read_file", {"path": str(root / "project" / f"f{n}.txt")})], input_tokens=usage) + return Turn([Reasoning("Every file has been read; write the final answer."), Text(FINAL_C)], + input_tokens=usage) + + return respond + + +def _scenario_compaction(root: Path) -> dict[str, Any]: + fake = FakeBedrock(_compaction_responder(root)) + with fake: + nh = _home(root, fake, extra_config={"compression": {"threshold_tokens": 12_000, "protect_last_n": 4}}) + for i in range(1, COMPACTION_FILES + 1): + (nh.project / f"f{i}.txt").write_text(f"file {i} " + "lorem ipsum dolor " * 250 + "\n", encoding="utf-8") + result = run_chat(nh, f"Read f1.txt through f{COMPACTION_FILES}.txt one by one, then say done.", + env=_endpoint_env(fake)) + return {"fake": fake, "nh": nh, "result": result, "requests": fake.snapshot()} + + +SCENARIOS: dict[str, Callable[[Path], dict[str, Any]]] = { + "tools": _scenario_tools, "resume": _scenario_resume, "compaction": _scenario_compaction, +} + + +@pytest.fixture(scope="module") +def runs(tmp_path_factory: pytest.TempPathFactory) -> dict[str, dict[str, Any]]: + base = tmp_path_factory.mktemp("bedrock_turns") + with ThreadPoolExecutor(len(SCENARIOS)) as pool: + futures = {name: pool.submit(fn, base / name) for name, fn in SCENARIOS.items()} + return {name: fut.result() for name, fut in futures.items()} + + +# -------------------------------------------------------------------------------------------------- +# Assertion helpers +# -------------------------------------------------------------------------------------------------- + + +def _ok(result: ChatResult) -> None: + assert result.returncode == 0, result.describe() + + +def _accepted(requests: list[dict[str, Any]]) -> None: + rejected = [r["rejected"] for r in requests if r.get("rejected")] + assert not rejected, f"Bedrock (fake) rejected {len(rejected)} request(s): {rejected}" + + +def _main(requests: list[dict[str, Any]]) -> list[dict[str, Any]]: + return [r for r in requests if "toolConfig" in (r.get("body") or {})] + + +def _results(message: dict[str, Any]) -> dict[str, str]: + return {b["toolResult"]["toolUseId"]: json.dumps(b["toolResult"]["content"]) + for b in message["content"] if "toolResult" in b} + + +def _tool_uses(blocks: list[dict[str, Any]]) -> list[dict[str, Any]]: + return [b["toolUse"] for b in blocks if "toolUse" in b] + + +def _blob(request: dict[str, Any]) -> str: + return json.dumps(request["body"].get("messages", [])) + + +# -------------------------------------------------------------------------------------------------- +# a. multi-turn with tools +# -------------------------------------------------------------------------------------------------- + + +def test_every_call_is_sigv4_signed_for_bedrock_with_the_profile_credentials(runs: dict[str, Any]) -> None: + requests = [r for run in runs.values() for r in run["requests"]] + assert requests, "no request reached the Bedrock fake" + _accepted(requests) + for rec in requests: + assert rec["auth"]["key"] == ACCESS_KEY and rec["auth"]["service"] == "bedrock", rec["headers"] + assert rec["auth"]["region"] == REGION, rec["auth"] + assert rec["model"] == MODEL and rec["op"] in ("Converse", "ConverseStream"), rec["path"] + + +def test_parallel_tool_uses_round_trip_as_tool_results_paired_by_id(runs: dict[str, Any]) -> None: + run = runs["tools"] + _ok(run["result"]) + requests = run["requests"] + _accepted(requests) + assert [r["op"] for r in requests] == ["ConverseStream"] * 3, [r["op"] for r in requests] + first_ids = [tu["toolUseId"] for tu in _tool_uses(requests[0]["emitted"])] + second = requests[1]["body"]["messages"] + # The assistant turn goes back byte-for-byte: signed reasoning first, then both toolUse blocks. + assert second[-2] == {"role": "assistant", "content": requests[0]["emitted"]}, second[-2] + results = _results(second[-1]) + assert list(results) == first_ids, f"toolResult ids {list(results)} != toolUse ids {first_ids}" + assert SEED_TEXT in results[first_ids[0]] and "MARK-42" in results[first_ids[1]], results + third = requests[2]["body"]["messages"] + assert third[:len(second)] == second, "history prefix changed between tool rounds" + assert third[-2] == {"role": "assistant", "content": requests[1]["emitted"]}, third[-2] + (second_id,) = [tu["toolUseId"] for tu in _tool_uses(requests[1]["emitted"])] + assert SEED2_TEXT in _results(third[-1])[second_id], third[-1] + + +def test_tool_turn_prints_final_answer_and_persists_paired_rows(runs: dict[str, Any]) -> None: + run = runs["tools"] + _ok(run["result"]) + assert FINAL_A in run["result"].stdout, run["result"].describe() + rows = messages(run["nh"]) + assert [r["role"] for r in rows] == ["user", "assistant", "tool", "tool", "assistant", "tool", "assistant"], rows + issued = [[tu["toolUseId"] for tu in _tool_uses(r["emitted"])] for r in run["requests"][:2]] + assert [tc["id"] for tc in tool_calls_of(rows[1])] == issued[0] + assert [r["tool_call_id"] for r in rows[2:4]] == issued[0] + assert [tc["id"] for tc in tool_calls_of(rows[4])] == issued[1] and rows[5]["tool_call_id"] == issued[1][0] + assert rows[-1]["content"] == FINAL_A + + +def _reasoning_rows(nh: NativeHome) -> list[str]: + return [r["reasoning_content"] for r in messages(nh) if r["role"] == "assistant" and r["reasoning_content"]] + + +@pytest.mark.xfail(strict=True, raises=KnownSymptom, reason=KNOWN["reasoning_shredded"]) +def test_persisted_reasoning_equals_the_streamed_reasoning_text(runs: dict[str, Any]) -> None: + _ok(runs["tools"]["result"]) + persisted = _reasoning_rows(runs["tools"]["nh"]) + assert [p.replace("\n\n", "") for p in persisted] == [R_A1, R_A2, R_A3], persisted + if persisted != [R_A1, R_A2, R_A3]: # same text once the blank lines are removed: exactly the bug + raise KnownSymptom(f"{KNOWN['reasoning_shredded']}: {persisted}") + + +# -------------------------------------------------------------------------------------------------- +# b. signed reasoning replay across --resume +# -------------------------------------------------------------------------------------------------- + + +def test_resume_in_new_process_replays_tool_history_valid_for_converse(runs: dict[str, Any]) -> None: + run = runs["resume"] + _ok(run["first"]) + _ok(run["second"]) + assert FINAL_B2 in run["second"].stdout, run["second"].describe() + requests = run["requests"] + _accepted(requests) + assert len(requests) == 3 and session_ids(run["nh"]) == [run["sid"]], (len(requests), session_ids(run["nh"])) + resumed = requests[2]["body"]["messages"] + (tool_use,) = _tool_uses(requests[0]["emitted"]) + assert [m["role"] for m in resumed] == ["user", "assistant", "user", "assistant", "user"], resumed + assert _tool_uses(resumed[1]["content"]) == [tool_use], resumed[1] + assert SEED_TEXT in _results(resumed[2])[tool_use["toolUseId"]], resumed[2] + assert {"text": FINAL_B1} in resumed[3]["content"] and resumed[4]["content"] == [{"text": "Say it again."}] + rows = messages(run["nh"], run["sid"]) + assert [r["role"] for r in rows] == ["user", "assistant", "tool", "assistant", "user", "assistant"], rows + assert rows[-1]["content"] == FINAL_B2 + + +@pytest.mark.xfail(strict=True, raises=KnownSymptom, reason=KNOWN["resume_drops_reasoning"]) +def test_resume_replays_signed_reasoning_verbatim(runs: dict[str, Any]) -> None: + run = runs["resume"] + _ok(run["first"]) + _ok(run["second"]) + requests = run["requests"] + assert len(requests) == 3, f"expected 2 turn-1 calls + 1 resumed call, fake saw {len(requests)}" + signed = [b for b in requests[0]["emitted"] if "reasoningContent" in b] + assert signed, requests[0]["emitted"] + # In-process the tool-use turn goes back with its signed reasoning ... + assert requests[1]["body"]["messages"][1]["content"][0] == signed[0] + # ... and after --resume (new process) it must be identical. + resumed = requests[2]["body"]["messages"] + if not [b for b in resumed[1]["content"] if "reasoningContent" in b]: + raise KnownSymptom(f"{KNOWN['resume_drops_reasoning']}: {resumed[1]['content']}") + assert resumed[1]["content"][0] == signed[0], resumed[1]["content"] + assert [b for b in resumed[3]["content"] if "reasoningContent" in b] == [ + b for b in requests[1]["emitted"] if "reasoningContent" in b] + + +# -------------------------------------------------------------------------------------------------- +# c. compaction in a reasoning session +# -------------------------------------------------------------------------------------------------- + + +def test_compaction_in_reasoning_session_keeps_every_request_converse_valid(runs: dict[str, Any]) -> None: + run = runs["compaction"] + _ok(run["result"]) + assert FINAL_C in run["result"].stdout, run["result"].describe() + requests = run["requests"] + _accepted(requests) # tool pairs, signatures and final-assistant thinking checked on every request + summaries = [i for i, r in enumerate(requests) if r["op"] == "Converse" and "toolConfig" not in r["body"]] + assert summaries, "compaction never called the summarizer over Converse" + after = [r for r in _main(requests[summaries[0] + 1:])] + assert after, "no main request after the compaction summary" + before_len = max(len(r["body"]["messages"]) for r in _main(requests[:summaries[0]])) + assert all(SUMMARY_TOKEN in _blob(r) for r in after), "a post-compaction request lost the summary" + assert len(after[0]["body"]["messages"]) < before_len, "compaction did not shrink the wire history" + last = after[-1]["body"]["messages"] + # The open tool loop's final assistant turn still leads with the signed reasoning Bedrock issued. + assert "reasoningContent" in last[-2]["content"][0] and _results(last[-1]), last[-2:] + + +def test_compaction_persists_summary_and_archives_compacted_rows(runs: dict[str, Any]) -> None: + run = runs["compaction"] + _ok(run["result"]) + every = messages(run["nh"], active_only=False) + live = messages(run["nh"]) + assert any(r["compacted"] == 1 for r in every), "no row archived as compacted" + assert any(SUMMARY_TOKEN in (r["content"] or "") for r in live), "summary not in the live transcript" + assert live[-1]["role"] == "assistant" and live[-1]["content"] == FINAL_C, live[-1] + live_ids = {r["tool_call_id"] for r in live if r["role"] == "tool"} + issued = {tc["id"] for r in live if r["role"] == "assistant" for tc in tool_calls_of(r)} + assert live_ids <= issued, f"orphaned live tool rows: {live_ids - issued}" + + +# -------------------------------------------------------------------------------------------------- +# The fake itself rejects what Bedrock rejects (so a green scenario above is not a pass-through) +# -------------------------------------------------------------------------------------------------- + + +def _mutate_request(case: str, kwargs: dict[str, Any]) -> dict[str, Any]: + msgs = kwargs["messages"] + table: dict[str, Callable[[], None]] = { + "orphan_tool_result": lambda: msgs[2]["content"][0]["toolResult"].update(toolUseId="tooluse_nobody"), + "tampered_signature": lambda: msgs[1]["content"][0]["reasoningContent"]["reasoningText"].update(signature="Zm9yZ2Vk"), + "unsigned_final_thinking": lambda: msgs[1]["content"].pop(0), + "two_members_in_union": lambda: msgs[0]["content"][0].update(image={"format": "png", "source": {"bytes": b"x"}}), + "blank_text": lambda: msgs[0]["content"][0].update(text=" "), + } + table[case]() + return kwargs + + +@pytest.mark.parametrize("case", ["valid", "orphan_tool_result", "tampered_signature", "unsigned_final_thinking", + "two_members_in_union", "blank_text", "bad_secret"]) +def test_fake_rejects_requests_bedrock_rejects(case: str) -> None: + boto3 = pytest.importorskip("boto3") + from botocore.config import Config + from botocore.exceptions import ClientError, ParamValidationError + + fake = FakeBedrock(seq(lambda _r: Turn([Reasoning("sig check reasoning"), ToolUse("noop", {})]), + Turn([Text("fine")]))) + with fake: + client = boto3.client( + "bedrock-runtime", region_name=REGION, endpoint_url=fake.endpoint, aws_access_key_id=ACCESS_KEY, + aws_secret_access_key="wrong-secret" if case == "bad_secret" else SECRET_KEY, + config=Config(retries={"max_attempts": 1}, parameter_validation=False)) + base = [{"role": "user", "content": [{"text": "go"}]}] + first = client.converse(modelId=MODEL, messages=base) if case != "bad_secret" else None + blocks = first["output"]["message"]["content"] if first else [] + tool_id = next((b["toolUse"]["toolUseId"] for b in blocks if "toolUse" in b), "none") + kwargs = {"modelId": MODEL, "messages": base + [ + {"role": "assistant", "content": blocks or [{"text": "x"}]}, + {"role": "user", "content": [{"toolResult": {"toolUseId": tool_id, "content": [{"text": "ok"}]}}]}]} + if case not in ("valid", "bad_secret"): + kwargs = _mutate_request(case, kwargs) + try: + client.converse(**kwargs) + except (ClientError, ParamValidationError) as exc: + code = getattr(exc, "response", {}).get("Error", {}).get("Code", type(exc).__name__) + else: + code = "OK" + expected = {"valid": "OK", "bad_secret": "InvalidSignatureException"}.get(case, "ValidationException") + assert code == expected, (case, code, fake.snapshot()[-1].get("rejected")) diff --git a/tests/e2e/core/providers/test_native_codex_app_server.py b/tests/e2e/core/providers/test_native_codex_app_server.py new file mode 100644 index 000000000000..6e6fdfd3f9b7 --- /dev/null +++ b/tests/e2e/core/providers/test_native_codex_app_server.py @@ -0,0 +1,206 @@ +"""codex_app_server wire conformance: real ``hermes chat -q`` against a fake ``codex app-server``. + +The fake (``tests/fakes/providers/codex_app_server.py``) speaks newline-delimited JSON-RPC over stdio, +validates every request/response Hermes sends against the codex-cli 0.147 app-server schema (serde-style +``-32600 Invalid request`` on missing/mistyped fields; unknown fields recorded because the real server +silently drops them) and records the transcript per app-server PID. Selected via +``model.openai_runtime: codex_app_server`` + ``model.codex_bin``. + +Happy-path lifecycle: tool item + approval round trip, reasoning projection, ``--resume`` in a new process +(``thread/resume`` vs. the history-seed fallback) and codex-native compaction. +""" + +from __future__ import annotations + +import json +import sys +from collections.abc import Iterator +from concurrent.futures import ThreadPoolExecutor + +import pytest + +from tests.e2e.core.providers._native_helpers import KnownSymptom, messages, tool_calls_of +from tests.fakes.providers.codex_app_server import CodexRun, run_codex_scenario + +pytestmark = [ + pytest.mark.skipif(sys.platform == "win32", reason="POSIX sh wrapper for the fake codex binary"), + # CodexRun.cleanup() SIGKILLs any app-server that outlived its CLI (reparented to init by then). + pytest.mark.live_system_guard_bypass, +] + +KNOWN = { + "compaction_row": "#121301 native contextCompaction item persisted as a raw-JSON assistant message", +} + +YOLO = ["--yolo"] +SEED_MARKER = "Prior conversation from this Hermes session" + +SCENARIOS = { + # a + b: reasoning, a command needing approval, usage, final answer; then --resume in a new process. + "tools": dict( + turns=[ + {"steps": [{"kind": "reasoning", "summary": ["R-SUMMARY-alpha weighing the canary"]}, + {"kind": "command", "command": "echo CANARY-1", "output": "CANARY-1-OUT\n"}, + {"kind": "usage", "input": 1234, "cached": 200, "output": 30}, + {"kind": "message", "text": "FINAL-ANSWER-ONE"}]}, + {"steps": [{"kind": "message", "text": "FINAL-ANSWER-TWO"}]}, + ], + runs=[{"prompt": "run the canary USER-ONE", "args": YOLO}, {"prompt": "and now USER-TWO", "args": YOLO}], + ), + # Resume fallback: the rollout is gone for run 2 (history seed on thread/start), back for run 3. + "seed": dict( + turns=[ + {"steps": [{"kind": "reasoning", "summary": ["R-SEED-PRIVATE never replayed"]}, + {"kind": "command", "command": "echo SEED-CMD", "output": "SEED-OUT\n"}, + {"kind": "message", "text": "SEED-FINAL-ONE"}]}, + {"steps": [{"kind": "message", "text": "SEED-FINAL-TWO"}]}, + {"steps": [{"kind": "message", "text": "SEED-FINAL-THREE"}]}, + ], + runs=[{"prompt": "first prompt USER-ONE", "args": YOLO, "then": {"forget_threads": True}}, + {"prompt": "second prompt USER-TWO", "args": YOLO, "then": {"forget_threads": False}}, + {"prompt": "third prompt USER-THREE", "args": YOLO}], + ), + # c: codex compacts natively mid-turn (contextCompaction item); the next process resumes the same thread. + "compact": dict( + turns=[ + {"steps": [{"kind": "message", "text": "C-ONE"}, {"kind": "compaction"}, + {"kind": "usage", "input": 500}, {"kind": "message", "text": "C-TWO"}]}, + {"steps": [{"kind": "message", "text": "C-THREE"}]}, + ], + runs=[{"prompt": "compact me", "args": YOLO}, {"prompt": "after compact", "args": YOLO}], + ), +} + + +@pytest.fixture(scope="module") +def runs(tmp_path_factory) -> Iterator[dict[str, CodexRun]]: + with ThreadPoolExecutor(max_workers=len(SCENARIOS)) as pool: + futures = {name: pool.submit(run_codex_scenario, tmp_path_factory.mktemp(f"codex_{name}"), **spec) + for name, spec in SCENARIOS.items()} + done = {name: future.result() for name, future in futures.items()} + yield done + for run in done.values(): + run.cleanup() + + +def _started_thread_id(run: CodexRun, index: int) -> str: + """Thread id the fake returned to a thread/start in the ``index``-th app-server process.""" + starts = {m["id"] for m in run.process_requests(index, "thread/start")} + ids = [e["msg"]["result"]["thread"]["id"] for e in run.process_entries(index) + if e.get("dir") == "out" and e["msg"].get("id") in starts and "result" in e["msg"]] + assert len(ids) == 1, f"expected one successful thread/start in process {index}, got {ids}" + return ids[0] + + +def _rows(run: CodexRun) -> list[dict]: + return messages(run.home, run.session_id) + + +def test_every_request_is_schema_valid_and_handshake_ordered(runs): + for name, run in runs.items(): + run.fake.assert_wire_clean() + for index, _ in enumerate(run.results): + inbound = [e["msg"] for e in run.process_entries(index) if e.get("dir") == "in"] + methods = [m.get("method") for m in inbound] + assert methods[:2] == ["initialize", "initialized"], f"{name}#{index} handshake order: {methods}" + assert "id" not in inbound[1], f"{name}#{index}: `initialized` must be a notification" + assert inbound[0]["params"]["clientInfo"]["name"], f"{name}#{index}: empty clientInfo.name" + + +def test_command_approval_round_trip_and_tool_rows(runs): + run = runs["tools"] + assert run.results[0].returncode == 0, run.results[0].describe() + assert run.results[0].stdout.count("FINAL-ANSWER-ONE") == 1, run.results[0].describe() + + first = run.process_entries(0) + request = next(e["msg"] for e in first if e.get("dir") == "out" + and e["msg"].get("method") == "item/commandExecution/requestApproval") + replies = [e for e in first if e.get("reply_to") == "item/commandExecution/requestApproval"] + assert [r["msg"]["id"] for r in replies] == [request["id"]], f"approval reply ids: {replies}" + assert replies[0]["msg"]["result"] == {"decision": "accept"}, replies[0] + + rows = _rows(run) + calls = [(row, call) for row in rows for call in tool_calls_of(row)] + assert len(calls) == 1, f"expected one projected tool call, rows: {rows}" + call_row, call = calls[0] + assert json.loads(call["function"]["arguments"])["command"] == "echo CANARY-1", call + results = [r for r in rows if r["role"] == "tool"] + assert [r["tool_call_id"] for r in results] == [call["id"]], f"tool result not paired: {results}" + assert "CANARY-1-OUT" in results[0]["content"], results[0] + finals = [r for r in rows if r["role"] == "assistant" and r["content"] == "FINAL-ANSWER-ONE"] + assert len(finals) == 1 and finals[0]["id"] > results[0]["id"], f"final answer row missing/out of order: {rows}" + + +def test_reasoning_projected_once_and_never_as_content(runs): + run = runs["tools"] + rows = _rows(run) + carriers = [r for r in rows if "R-SUMMARY-alpha" in (r.get("reasoning") or "")] + assert len(carriers) == 1, f"reasoning must attach to exactly one row (after resume too): {rows}" + assert carriers[0]["role"] == "assistant" and tool_calls_of(carriers[0]), carriers[0] + assert not [r for r in rows if "R-SUMMARY-alpha" in (r.get("content") or "")], "reasoning leaked into content" + assert "R-SUMMARY-alpha" not in run.output, "reasoning printed in quiet mode" + + +def test_resume_in_new_process_resumes_the_same_thread(runs): + run = runs["tools"] + assert run.results[1].returncode == 0 and "FINAL-ANSWER-TWO" in run.results[1].stdout, run.results[1].describe() + thread_id = _started_thread_id(run, 0) + assert run.fake.spawned_pids()[0] != run.fake.spawned_pids()[1] + assert run.process_requests(1, "thread/start") == [], "resume must not start a fresh thread" + resumes = run.process_requests(1, "thread/resume") + assert [r["params"]["threadId"] for r in resumes] == [thread_id], resumes + assert SEED_MARKER not in (resumes[0]["params"].get("developerInstructions") or ""), \ + "a resumed thread already holds the history; seeding it again duplicates the conversation" + turn_starts = run.process_requests(1, "turn/start") + assert [(t["params"]["threadId"], t["params"]["input"]) for t in turn_starts] == \ + [(thread_id, [{"type": "text", "text": "and now USER-TWO"}])], turn_starts + contents = [(r["role"], r["content"]) for r in _rows(run) if r["content"]] + assert contents.count(("assistant", "FINAL-ANSWER-ONE")) == 1 and contents[-2:] == [ + ("user", "and now USER-TWO"), ("assistant", "FINAL-ANSWER-TWO")], contents + + +def test_resume_fallback_seeds_history_then_rebinds_new_thread(runs): + run = runs["seed"] + assert [r.returncode for r in run.results] == [0, 0, 0], "\n".join(r.describe() for r in run.results) + old_thread = _started_thread_id(run, 0) + assert [r["params"]["threadId"] for r in run.process_requests(1, "thread/resume")] == [old_thread] + new_thread = _started_thread_id(run, 1) + seed = run.process_requests(1, "thread/start")[0]["params"].get("developerInstructions") or "" + assert SEED_MARKER in seed, "fallback thread/start must carry the prior conversation" + history = seed[seed.index(SEED_MARKER):] + for needle in ("first prompt USER-ONE", "SEED-OUT", "SEED-FINAL-ONE"): + assert history.count(needle) == 1, f"{needle!r} seeded {history.count(needle)}x:\n{history}" + assert "R-SEED-PRIVATE" not in history, "reasoning must not be replayed as conversation text" + assert "second prompt USER-TWO" not in history, "the new user turn goes in turn/start, not the seed" + assert [t["params"]["threadId"] for t in run.process_requests(1, "turn/start")] == [new_thread] + # Run 3: the session is now bound to the replacement thread. + assert [r["params"]["threadId"] for r in run.process_requests(2, "thread/resume")] == [new_thread] + assert run.process_requests(2, "thread/start") == [] + contents = [r["content"] for r in _rows(run) if r["role"] == "assistant" and r["content"]] + assert contents == ["SEED-FINAL-ONE", "SEED-FINAL-TWO", "SEED-FINAL-THREE"], contents + + +def test_native_compaction_keeps_thread_and_transcript(runs): + run = runs["compact"] + assert [r.returncode for r in run.results] == [0, 0], "\n".join(r.describe() for r in run.results) + assert "C-TWO" in run.results[0].stdout and "C-THREE" in run.results[1].stdout + thread_id = _started_thread_id(run, 0) + assert [r["params"]["threadId"] for r in run.process_requests(1, "thread/resume")] == [thread_id], \ + "codex-native compaction must not retire the thread" + assert run.process_requests(1, "thread/start") == [] + assert run.fake.requests("thread/compact/start") == [], "native mode: Hermes must not compact on top of codex" + rows = _rows(run) + assert {r["session_id"] for r in messages(run.home)} == {run.session_id}, "session was split" + texts = [r["content"] for r in rows if r["role"] == "assistant" and r["content"] in ("C-ONE", "C-TWO", "C-THREE")] + assert texts == ["C-ONE", "C-TWO", "C-THREE"], f"transcript rewritten or duplicated: {rows}" + + +@pytest.mark.xfail(strict=True, raises=KnownSymptom, reason=KNOWN["compaction_row"]) +def test_native_compaction_is_not_persisted_as_assistant_text(runs): + run = runs["compact"] + assert [r.returncode for r in run.results] == [0, 0], "\n".join(r.describe() for r in run.results) + rows = _rows(run) + assert any(r["content"] == "C-TWO" for r in rows), f"post-compaction answer not persisted: {rows}" + leaked = [r["content"] for r in rows if "contextCompaction" in (r.get("content") or "")] + if leaked: + raise KnownSymptom(f"compaction boundary persisted as assistant content: {leaked}") diff --git a/tests/e2e/core/providers/test_native_codex_app_server_faults.py b/tests/e2e/core/providers/test_native_codex_app_server_faults.py new file mode 100644 index 000000000000..0629e06f25b1 --- /dev/null +++ b/tests/e2e/core/providers/test_native_codex_app_server_faults.py @@ -0,0 +1,147 @@ +"""codex_app_server faults and approval edges: real ``hermes chat -q`` against a fake ``codex app-server``. + +See ``test_native_codex_app_server.py`` for the fake. Here: the app-server crashing mid-item, a failed turn, +a JSON-RPC error on ``turn/start``, a retrying ``error`` notification, and the server-initiated requests +Hermes must answer (approval in single-query mode, permissions) plus process-tree teardown. +""" + +from __future__ import annotations + +import sys +from collections.abc import Iterator +from concurrent.futures import ThreadPoolExecutor + +import pytest + +from tests.e2e.core.providers._native_helpers import KnownSymptom, messages, wait_until +from tests.fakes.providers.codex_app_server import CodexRun, pid_alive, run_codex_scenario + +pytestmark = [ + pytest.mark.skipif(sys.platform == "win32", reason="POSIX sh wrapper + /proc PID checks"), + # The orphaned own-session grandchild (#121298) is released cooperatively by CodexRun.cleanup(); the + # bypass covers its SIGKILL fallback for anything still alive (reparented to init by then). + pytest.mark.live_system_guard_bypass, +] + +KNOWN = { + "q_approval": "#121296 approval in `chat -q` waits the full approvals.timeout instead of single_query_mode", + "permissions": "#121297 reply to item/permissions/requestApproval omits required `permissions`", + "orphan": "#121298 `chat -q` exit never closes the codex session; own-session descendants orphaned", + "failed_hidden": "#121299 failed turn after an agentMessage prints the message and hides the reason", +} + +YOLO = ["--yolo"] +APPROVAL_TIMEOUT_S = 8 + + +def _one(turn: dict, *, args=YOLO, config=None) -> dict: + return dict(turns=[turn], runs=[{"prompt": "go", "args": args}], config=config) + + +SCENARIOS = { + "crash": _one({"steps": [{"kind": "message_partial", "text": "PARTIAL-A-CRASH"}, + {"kind": "crash", "code": 3, "stderr": "fatal: CRASH-MARKER-77"}]}), + "failed": _one({"steps": [{"kind": "fail", "message": "FAIL-MARKER-88 stream disconnected before completion"}]}), + "start_error": _one({"start_error": "START-ERR-99 model is overloaded"}), + "retry_note": _one({"steps": [{"kind": "error_note", "message": "RECONNECT-NOTE 1/5", "will_retry": True}, + {"kind": "message", "text": "RETRY-OK"}]}), + "failed_hidden": _one({"steps": [{"kind": "message", "text": "PARTIAL-B"}, + {"kind": "fail", "message": "FAIL-MARKER-89 stream disconnected"}]}), + "q_approval": _one({"steps": [{"kind": "command", "command": "echo Q", "output": "Q\n"}, + {"kind": "message", "text": "Q-DONE"}]}, + args=[], config={"approvals": {"timeout": APPROVAL_TIMEOUT_S}}), + "permissions": _one({"steps": [{"kind": "permissions", "reason": "needs network"}, + {"kind": "message", "text": "PERM-DONE"}]}), + "orphan": _one({"steps": [{"kind": "grandchild"}, {"kind": "message", "text": "REAP-DONE"}]}), +} + + +@pytest.fixture(scope="module") +def runs(tmp_path_factory) -> Iterator[dict[str, CodexRun]]: + with ThreadPoolExecutor(max_workers=len(SCENARIOS)) as pool: + futures = {name: pool.submit(run_codex_scenario, tmp_path_factory.mktemp(f"codex_{name}"), **spec) + for name, spec in SCENARIOS.items()} + done = {name: future.result() for name, future in futures.items()} + yield done + for run in done.values(): + run.cleanup() + + +def _assistant_texts(run: CodexRun) -> list[str]: + return [r["content"] for r in messages(run.home, run.session_id) if r["role"] == "assistant" and r["content"]] + + +def test_crash_mid_item_surfaced_once_without_partial_content_or_orphan(runs): + run = runs["crash"] + result = run.results[0] + assert result.returncode != 0, result.describe() + assert run.output.count("exited unexpectedly") == 1, result.describe() + assert run.output.count("CRASH-MARKER-77") == 1, "the app-server's stderr tail must reach the user once" + assert "PARTIAL-A-CRASH" not in run.output, "an unfinished item must not be shown as the answer" + assert not [t for t in _assistant_texts(run) if "PARTIAL-A-CRASH" in t], "unfinished item persisted" + assert len(run.fake.spawned_pids()) == 1 and len(run.fake.requests("turn/start")) == 1, \ + "a crashed turn must not be replayed on a respawned app-server" + assert not pid_alive(run.fake.spawned_pids()[0]), "app-server process survived the CLI" + + +@pytest.mark.parametrize("name, marker", [("failed", "FAIL-MARKER-88"), ("start_error", "START-ERR-99")]) +def test_terminal_turn_error_surfaced_once_not_retried(runs, name, marker): + run = runs[name] + assert run.results[0].returncode != 0, run.results[0].describe() + assert run.output.count(marker) == 1, run.results[0].describe() + assert len(run.fake.requests("turn/start")) == 1, "non-retryable turn error was retried" + assert not [t for t in _assistant_texts(run) if marker in t], "error text persisted as assistant content" + run.fake.assert_wire_clean() + + +def test_will_retry_error_notification_is_not_terminal(runs): + run = runs["retry_note"] + assert run.results[0].returncode == 0, run.results[0].describe() + assert run.results[0].stdout.strip().endswith("RETRY-OK"), run.results[0].describe() + assert "RECONNECT-NOTE" not in run.output, "a willRetry notice is codex's own retry, not a user-facing error" + assert len(run.fake.requests("turn/start")) == 1 + assert _assistant_texts(run) == ["RETRY-OK"] + + +@pytest.mark.xfail(strict=True, raises=KnownSymptom, reason=KNOWN["failed_hidden"]) +def test_failed_turn_after_agent_message_surfaces_reason(runs): + run = runs["failed_hidden"] + assert run.results[0].returncode != 0 and "PARTIAL-B" in run.results[0].stdout, run.results[0].describe() + if "FAIL-MARKER-89" not in run.output: + raise KnownSymptom(f"turn failure reason never shown to the user: {run.output!r}") + + +@pytest.mark.xfail(strict=True, raises=KnownSymptom, reason=KNOWN["q_approval"]) +def test_single_query_approval_resolves_without_waiting_for_a_human(runs): + run = runs["q_approval"] + entries = run.fake.entries() + sent = [e for e in entries if e.get("dir") == "out" + and e["msg"].get("method") == "item/commandExecution/requestApproval"] + replies = [e for e in entries if e.get("reply_to") == "item/commandExecution/requestApproval"] + assert len(sent) == 1 and len(replies) == 1, f"approval request/reply not exchanged once: {sent} {replies}" + reply = replies[0] + assert reply["msg"]["id"] == sent[0]["msg"]["id"] and not reply.get("violation"), reply + waited = reply["t"] - sent[0]["t"] + if waited >= APPROVAL_TIMEOUT_S / 2: + raise KnownSymptom(f"approval parked {waited:.1f}s on a prompt nobody can answer in -q") + + +@pytest.mark.xfail(strict=True, raises=KnownSymptom, reason=KNOWN["permissions"]) +def test_permissions_request_reply_matches_protocol(runs): + run = runs["permissions"] + replies = run.fake.replies_to("item/permissions/requestApproval") + assert len(replies) == 1, f"permissions request unanswered: {replies}" + violation = replies[0].get("violation") or "" + if "permissions" in violation: + raise KnownSymptom(f"PermissionsRequestApprovalResponse without `permissions`: {replies[0]}") + assert not violation, f"invalid PermissionsRequestApprovalResponse: {replies[0]}" + + +@pytest.mark.xfail(strict=True, raises=KnownSymptom, reason=KNOWN["orphan"]) +def test_cli_exit_reaps_app_server_descendants(runs): + run = runs["orphan"] + assert run.results[0].returncode == 0 and "REAP-DONE" in run.results[0].stdout, run.results[0].describe() + children = run.fake.grandchild_pids() + assert len(children) == 1, f"the fake must have spawned exactly one descendant: {children}" + wait_until(lambda: not pid_alive(children[0]), 3.0, + f"app-server descendant {children[0]} to be reaped after CLI exit", error=KnownSymptom) diff --git a/tests/e2e/core/providers/test_native_copilot_acp.py b/tests/e2e/core/providers/test_native_copilot_acp.py new file mode 100644 index 000000000000..ac0bdbbf62ea --- /dev/null +++ b/tests/e2e/core/providers/test_native_copilot_acp.py @@ -0,0 +1,313 @@ +"""Copilot ACP wire conformance, part 1: the multi-call tool flow, resume, and the process lifecycle. + +``provider: copilot-acp`` makes Hermes spawn an external ACP agent (``copilot --acp --stdio``) per model +call and speak the Agent Client Protocol to it over stdio. Here the agent is +``tests/fakes/providers/copilot_acp.py``, a fake that validates every request against the published +ACP schema and replays scripted turns. Everything on the Hermes side is real: the ``hermes chat -q`` +process, runtime resolution from ``config.yaml`` + the profile ``.env`` (``HERMES_COPILOT_ACP_COMMAND`` +/ ``HERMES_COPILOT_ACP_ARGS``), the ACP client, the agent loop, the ``read_file`` tool and ``state.db``. + +Contract under test (documented in ``agent/copilot_acp_client.py`` and the ACP spec): + +* each model call is a fresh agent process: ``initialize`` -> ``session/new`` (absolute cwd) -> + model selection via the advertised ``model`` config option -> ``session/prompt``; every request is + schema-valid; +* ACP has no tools channel: Hermes' tools travel in the prompt text, a ```` block in the + agent's message runs a REAL Hermes tool, and the result is in the next call's prompt; +* ``--resume`` in a new process runs in a new agent process whose prompt carries the persisted history + (turn 1 in order, then the new question), with nothing duplicated; +* agent-side ``session/request_permission`` is never granted (Hermes has no human channel there) and + ``fs/read_text_file`` is confined to the session cwd; +* no agent process outlives the CLI, including one that ignores SIGTERM and stdin EOF. +""" + +from __future__ import annotations + +import contextlib +import os +import signal +import sys +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, Callable + +import pytest + +from tests.e2e.core.providers._native_helpers import ( + ChatResult, + KnownSymptom, + NativeHome, + assert_no_duplicate_assistant_text, + latest_session, + make_home, + messages, + run_chat, + session_ids, + tool_calls_of, + wait_until, +) +from tests.fakes.providers import copilot_acp as acp + +pytest.importorskip("acp.schema", reason="the fake validates against the agent-client-protocol package (acp extra)") + +pytestmark = [ + pytest.mark.skipif(not sys.platform.startswith("linux"), reason="POSIX launcher script + /proc pid checks"), + pytest.mark.live_system_guard_bypass, +] + +MODELS = ["fake-model-a", "fake-model-b"] +CONFIGURED_MODEL = "fake-model-b" # NOT the agent's default: selection must happen on the wire +CANARY = "CANARY-7731-acp" +SECRET = "OUTSIDE-CWD-SECRET-5512" +Q1 = "Read canary.txt and tell me what it says (Q1-marker)." +Q2 = "Now summarise what you found (Q2-marker)." +FINAL_ONE = f"The file says {CANARY} (FINAL-ONE)" +FINAL_TWO = "Summary: the canary was read (FINAL-TWO)" +LATE_TEXT = "LATE-ANSWER-65788" +# Compaction: the system prompt + tool bridge alone is ~18K estimated tokens; eight ~1.2K-token file reads +# cross this absolute threshold mid-turn (ACP reports no usage, so Hermes estimates). +COMPACT_THRESHOLD = 21_000 +COMPACT_FILES = 8 +COMPACT_ASK = "Read f1.txt through f8.txt one by one, then say done (COMPACT-ASK)." +SUMMARY = "SUMMARY-ACP-7f3: files f1..fN were read; each is lorem ipsum filler." +FINAL_COMPACT = "All eight files read (FINAL-COMPACT)." + + +KNOWN: dict[str, str] = { + "late_chunk": "#65788 agent_message_chunk emitted after the session/prompt result is dropped", +} + + +@dataclass +class Scenario: + nh: NativeHome + fake: acp.AcpFake + runs: list[ChatResult] = field(default_factory=list) + extra: dict[str, Any] = field(default_factory=dict) + + +def _home(root: Path, turns: list[list[dict[str, Any]]], **fake_kw: Any) -> Scenario: + fake = acp.AcpFake(root / "acp", turns, models=MODELS, **fake_kw) + nh = make_home(root, {"provider": "copilot-acp", "default": CONFIGURED_MODEL}, env_file=fake.env()) + return Scenario(nh, fake) + + +def _alive(pid: int) -> bool: + """True while ``pid`` is a live (non-zombie) process.""" + try: + with open(f"/proc/{pid}/stat", encoding="utf-8") as fh: + return fh.read().rsplit(")", 1)[1].split()[0] != "Z" + except (FileNotFoundError, ProcessLookupError, IndexError): + return False + + +def _reap(fake: acp.AcpFake) -> None: + """Teardown: SIGKILL any fake agent this scenario spawned that is still alive (by recorded pid).""" + for pid in fake.pids(): + if _alive(pid): + with contextlib.suppress(ProcessLookupError): # exited between the check and the kill + os.kill(pid, signal.SIGKILL) + + +# ── scenarios (independent; run concurrently in the module fixture) ────────────────────────── + + +def _flow(root: Path) -> Scenario: + """Turn 1 (two model calls): agent-side permission + a Hermes read_file call, then fs reads + the + answer. Turn 2: ``--resume`` in a new process. The agent ignores SIGTERM and stdin EOF.""" + project = NativeHome(root).project + turns = [ + [acp.tool_call("agent-native-1", "run shell", kind="execute"), acp.permission("agent-native-1"), + acp.message(acp.hermes_tool_call("call_rf_1", "read_file", {"path": str(project / "canary.txt")}))], + [acp.fs_read(str(project / "canary.txt")), acp.fs_read(str(root / "outside" / "secret.txt")), + acp.thought("REASONING-ONE: the tool result names the canary"), acp.message(FINAL_ONE)], + [acp.message(FINAL_TWO)], + ] + sc = _home(root, turns, ignore_sigterm=True) + (sc.nh.project / "canary.txt").write_text(CANARY + "\n", encoding="utf-8") + (root / "outside").mkdir() + (root / "outside" / "secret.txt").write_text(SECRET + "\n", encoding="utf-8") + sc.runs.append(run_chat(sc.nh, Q1)) + if sc.runs[0].returncode == 0 and session_ids(sc.nh): + sc.extra["turn1_pids"] = sc.fake.pids() + sc.runs.append(run_chat(sc.nh, Q2, resume=latest_session(sc.nh))) + return sc + + +def _late(root: Path) -> Scenario: + """The agent answers ``session/prompt`` first and streams the message chunk 250 ms later.""" + sc = _home(root, [[acp.result(), acp.message(LATE_TEXT, delay=0.25)], [acp.message("SECOND-TRY-65788")]]) + sc.runs.append(run_chat(sc.nh, "Say the late answer.")) + return sc + + +def _compaction(root: Path) -> Scenario: + """Eight read_file calls in one turn cross ``compression.threshold_tokens``; the summarizer runs + through the same ACP provider (no tool bridge -> the fake answers ``SUMMARY``).""" + project = NativeHome(root).project + turns = [[acp.thought(f"step {i}"), acp.message(acp.hermes_tool_call( + f"call_f{i}", "read_file", {"path": str(project / f"f{i}.txt")}))] for i in range(1, COMPACT_FILES + 1)] + fake = acp.AcpFake(root / "acp", [*turns, [acp.message(FINAL_COMPACT)]], models=MODELS, aux_text=SUMMARY) + nh = make_home(root, {"provider": "copilot-acp", "default": CONFIGURED_MODEL}, env_file=fake.env(), + extra_config={"compression": {"threshold_tokens": COMPACT_THRESHOLD, "protect_last_n": 4}}) + for i in range(1, COMPACT_FILES + 1): + (nh.project / f"f{i}.txt").write_text(f"file {i} " + "lorem ipsum dolor " * 250 + "\n", encoding="utf-8") + sc = Scenario(nh, fake) + sc.runs.append(run_chat(nh, COMPACT_ASK)) + return sc + + +SCENARIOS: dict[str, Callable[[Path], Scenario]] = {"flow": _flow, "late": _late, "compaction": _compaction} + + +@pytest.fixture(scope="module") +def outcomes(tmp_path_factory: pytest.TempPathFactory): + base = tmp_path_factory.mktemp("copilot_acp") + with ThreadPoolExecutor(max_workers=len(SCENARIOS)) as pool: + futures = {name: pool.submit(fn, base / name) for name, fn in SCENARIOS.items()} + done = {name: fut.result() for name, fut in futures.items()} + yield done + for sc in done.values(): + _reap(sc.fake) + + +def _flow_ok(outcomes: dict[str, Scenario]) -> Scenario: + sc = outcomes["flow"] + assert len(sc.runs) == 2 and all(r.returncode == 0 for r in sc.runs), "\n\n".join(r.describe() for r in sc.runs) + return sc + + +def _calls(fake: acp.AcpFake) -> dict[int, list[dict[str, Any]]]: + """Inbound client requests grouped per agent process (pid), in arrival order.""" + grouped: dict[int, list[dict[str, Any]]] = {} + for rec in fake.inbound(): + if rec["msg"] and rec["msg"].get("method"): + grouped.setdefault(rec["pid"], []).append(rec) + return grouped + + +# ── tests ────────────────────────────────────────────────────────────────────────────────────── + + +def test_hermes_tool_call_round_trips_through_acp_and_persists(outcomes): + """A ```` in the agent's message runs Hermes' real read_file; the result reaches the + NEXT call's prompt; the CLI prints the answer; state.db pairs the call and result by id.""" + sc = _flow_ok(outcomes) + assert FINAL_ONE in sc.runs[0].stdout, sc.runs[0].describe() + prompts = sc.fake.main_prompts() + assert len(prompts) == 3, f"expected 3 model calls (tool, answer, resumed answer), got {len(prompts)}" + first, second = acp.prompt_text(prompts[0]), acp.prompt_text(prompts[1]) + assert '"name": "read_file"' in first, "read_file schema was not offered through the prompt tool bridge" + assert CANARY not in first, "the canary leaked into the prompt before the tool ran" + assert CANARY in second and second.index(Q1) < second.index(CANARY), ( + "the read_file result must follow the user turn in the next call's prompt") + + rows = messages(sc.nh, latest_session(sc.nh)) + calls = [(r, tc) for r in rows if r["role"] == "assistant" for tc in tool_calls_of(r)] + assert [(tc["id"], tc["function"]["name"]) for _, tc in calls] == [("call_rf_1", "read_file")], calls + results = [r for r in rows if r["role"] == "tool"] + assert [r["tool_call_id"] for r in results] == ["call_rf_1"] and CANARY in results[0]["content"], results + finals = [r for r in rows if r["role"] == "assistant" and FINAL_ONE in (r["content"] or "")] + assert len(finals) == 1 and rows.index(finals[0]) > rows.index(results[0]), rows + assert "" not in "".join(r["content"] or "" for r in rows if r["role"] == "assistant"), ( + "raw tool-call bridge markup was persisted as assistant text") + assert "REASONING-ONE" in (finals[0].get("reasoning") or ""), "agent_thought_chunk text was not kept as reasoning" + + +def test_every_request_is_schema_valid_and_selects_the_configured_model(outcomes): + """Per process: initialize -> session/new (absolute project cwd) -> set_config_option(model) -> + session/prompt on the issued sessionId; zero requests the fake had to reject.""" + sc = _flow_ok(outcomes) + assert sc.fake.invalid() == [], f"requests rejected by the ACP schema: {sc.fake.invalid()}" + grouped = _calls(sc.fake) + assert len(grouped) == 3, f"one agent process per model call expected, got {len(grouped)}" + for pid, recs in grouped.items(): + methods = [r["msg"]["method"] for r in recs if r["msg"]["method"] not in ("session/cancel",)] + assert methods == ["initialize", "session/new", "session/set_config_option", "session/prompt"], (pid, methods) + init, new, select, prompt = (r["msg"]["params"] for r in recs[:4]) + assert init["protocolVersion"] == acp.PROTOCOL_VERSION + assert Path(new["cwd"]) == sc.nh.project.resolve(), new + assert select["value"] == CONFIGURED_MODEL and select["configId"] == "model", select + assert select["sessionId"] == prompt["sessionId"], (select, prompt) + + +def test_agent_permission_is_never_granted_and_fs_reads_stay_in_cwd(outcomes): + """The agent's own permission request is refused; fs/read_text_file inside the session cwd returns + the file, outside it returns a JSON-RPC error and no content.""" + sc = _flow_ok(outcomes) + outcomes_seen = [r["msg"]["result"]["outcome"] for r in sc.fake.records() if r.get("kind") == "permission_outcome"] + assert len(outcomes_seen) == 1, sc.fake.records() + assert outcomes_seen[0].get("outcome") != "selected", f"Hermes granted an agent-side permission: {outcomes_seen}" + reads = [r for r in sc.fake.records() if r.get("kind") == "fs_read_result"] + assert len(reads) == 2 and not any(r["errors"] for r in reads), reads + inside, outside = reads[0]["msg"], reads[1]["msg"] + assert CANARY in inside["result"]["content"], inside + assert "error" in outside and SECRET not in str(outside), f"fs read escaped the session cwd: {outside}" + + +def test_resume_reaches_the_agent_with_persisted_history_in_order(outcomes): + """``--resume`` in a new CLI process: a new agent process whose prompt carries turn 1 (question, tool + result, answer) in order and then the new question; one session, nothing persisted twice. How the + ACP session is opened (seeded ``session/new`` or ``session/load``) is not part of the contract.""" + sc = _flow_ok(outcomes) + assert FINAL_TWO in sc.runs[1].stdout, sc.runs[1].describe() + resumed_pids = [pid for pid in _calls(sc.fake) if pid not in sc.extra["turn1_pids"]] + assert len(resumed_pids) == 1, "the resumed turn must run in its own agent process" + text = acp.prompt_text(sc.fake.main_prompts()[-1]) + order = [text.find(s) for s in (Q1, CANARY, FINAL_ONE, Q2)] + assert -1 not in order and order == sorted(order), f"resumed prompt lost or reordered history: {order}" + sessions = session_ids(sc.nh) + rows = messages(sc.nh, sessions[-1]) + assert len(sessions) == 1 and [r["content"] for r in rows if r["role"] == "user"] == [Q1, Q2], rows + assert_no_duplicate_assistant_text(rows, FINAL_ONE) + assert_no_duplicate_assistant_text(rows, FINAL_TWO) + + +def test_no_agent_process_outlives_the_cli(outcomes): + """Every spawned agent is gone once the CLI exits, even one ignoring SIGTERM and stdin EOF.""" + sc = _flow_ok(outcomes) + pids = sc.fake.pids() + wedged = {r["pid"] for r in sc.fake.events("signal")} + assert len(pids) == 3 and wedged, "vacuity: the scenario must spawn 3 agents that ignored SIGTERM" + wait_until(lambda: not [p for p in pids if _alive(p)], 10.0, f"agent processes {pids} to exit") + + +def _transcript(record: dict[str, Any]) -> str: + text = acp.prompt_text(record) + return text[text.find("Conversation transcript:"):] + + +def test_compaction_in_an_acp_session_keeps_the_next_prompt_valid_and_grounded(outcomes): + """Auto compaction mid-turn: the summary is produced through the ACP agent itself, the next + main-turn prompt is schema-valid and carries the summary + the user's ask + the protected tail + (latest tool result) while the summarized tool output is gone; the turn completes once.""" + sc = outcomes["compaction"] + run = sc.runs[0] + assert run.returncode == 0 and FINAL_COMPACT in run.stdout, run.describe() + assert sc.fake.invalid() == [], f"requests rejected by the ACP schema: {sc.fake.invalid()}" + aux = sc.fake.aux_prompts() + assert aux, "compaction never called the summarizer through the ACP provider" + assert "file 2 lorem" in acp.prompt_text(aux[0]), "the summarizer did not receive the history to compact" + after = [r for r in sc.fake.main_prompts() if r["t"] > aux[0]["t"]] + assert after, "no main-turn call followed the compaction" + final = _transcript(after[-1]) + assert SUMMARY in final and COMPACT_ASK in final, "post-compaction prompt lost the summary or the user's ask" + assert f"file {COMPACT_FILES} lorem" in final, "post-compaction prompt lost the latest tool result" + assert "file 2 lorem" not in final, "summarized tool output is still resent after compaction" + rows = messages(sc.nh, latest_session(sc.nh)) + assert_no_duplicate_assistant_text(rows, FINAL_COMPACT) + assert any(FINAL_COMPACT in (r["content"] or "") for r in rows if r["role"] == "assistant"), rows + + +@pytest.mark.xfail(strict=True, raises=KnownSymptom, reason=KNOWN["late_chunk"]) +def test_message_chunk_after_prompt_result_reaches_the_user(outcomes): + sc = outcomes["late"] + run = sc.runs[0] + assert run.returncode == 0, run.describe() + assert sc.fake.invalid() == [], sc.fake.invalid() + rows = messages(sc.nh, latest_session(sc.nh)) + if LATE_TEXT not in run.stdout: + raise KnownSymptom(f"late chunk dropped: {len(sc.fake.main_prompts())} model calls, stdout={run.stdout!r}") + assert_no_duplicate_assistant_text(rows, LATE_TEXT) + assert len(sc.fake.main_prompts()) == 1, "a delivered late chunk must not trigger an empty-response retry" diff --git a/tests/e2e/core/providers/test_native_copilot_acp_errors.py b/tests/e2e/core/providers/test_native_copilot_acp_errors.py new file mode 100644 index 000000000000..62a651505000 --- /dev/null +++ b/tests/e2e/core/providers/test_native_copilot_acp_errors.py @@ -0,0 +1,162 @@ +"""Copilot ACP wire conformance, part 2: ACP error responses and agent crashes. + +The agent (``tests/fakes/providers/copilot_acp.py``) answers ``session/prompt`` with JSON-RPC errors from +the ACP / JSON-RPC 2.0 error space or dies mid-stream. Each row drives one real ``hermes chat -q`` and +asserts the retry semantics the user sees: + +* transient failures (``-32603`` internal error, a crash after a partial chunk) are retried with a + fresh agent process and the turn then succeeds, with none of the failed attempt's text persisted; +* ``-32000 Authentication required`` is not retried: surfaced once, one model call; +* persistent failures (``-32602`` invalid params, repeated crashes) stop after the configured retry + budget (``agent.api_max_retries: 2``) and are surfaced ONCE — never a loop, never a duplicate; +* no agent process survives the CLI in any row. +""" + +from __future__ import annotations + +import contextlib +import os +import re +import signal +import subprocess +import sys +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any + +import pytest + +from tests.e2e.core._pending_fixes import known_failure +from tests.e2e.core.providers._native_helpers import ( + ChatResult, + KnownSymptom, + NativeHome, + latest_session, + make_home, + messages, + run_chat, + wait_until, +) +from tests.fakes.providers import copilot_acp as acp + +pytest.importorskip("acp.schema", reason="the fake validates against the agent-client-protocol package (acp extra)") + +pytestmark = [ + pytest.mark.skipif(not sys.platform.startswith("linux"), reason="POSIX launcher script + /proc pid checks"), + pytest.mark.live_system_guard_bypass, +] + +API_MAX_RETRIES = 2 # _native_helpers.make_home pins agent.api_max_retries +OK_TEXT = "RECOVERED-ANSWER-OK" +PARTIAL = "PARTIAL-BEFORE-CRASH " +CRASH_STDERR = "fatal: agent segfaulted (fake)" + + +KNOWN: dict[str, str] = { + "auth_remedy": "#121290 copilot-acp auth failure tells the user to run a hermes command that is not implemented", +} + + +@dataclass(frozen=True) +class Row: + turns: list[list[dict[str, Any]]] + succeeds: bool + model_calls: int + visible: str # text the user must see exactly once (answer or the agent's error message) + never_persisted: tuple[str, ...] = () + + +ROWS: dict[str, Row] = { + "internal_error_retried": Row( + [[acp.rpc_error(-32603, "Internal error: upstream hiccup")], [acp.message(OK_TEXT)]], + succeeds=True, model_calls=2, visible=OK_TEXT, never_persisted=("upstream hiccup",)), + "crash_mid_stream_retried": Row( + [[acp.message(PARTIAL), acp.crash(3, CRASH_STDERR)], [acp.message(OK_TEXT)]], + succeeds=True, model_calls=2, visible=OK_TEXT, never_persisted=(PARTIAL.strip(),)), + "auth_required_not_retried": Row( + [[acp.rpc_error(-32000, "Authentication required")]] * 3, + succeeds=False, model_calls=1, visible="Authentication required"), + "invalid_params_bounded": Row( + [[acp.rpc_error(-32602, "Invalid params: prompt exceeds agent limit")]] * 4, + succeeds=False, model_calls=API_MAX_RETRIES, visible="prompt exceeds agent limit"), + "crash_every_time_bounded": Row( + [[acp.message(PARTIAL), acp.crash(3, CRASH_STDERR)]] * 4, + succeeds=False, model_calls=API_MAX_RETRIES, visible=CRASH_STDERR, never_persisted=(PARTIAL.strip(),)), +} + + +@dataclass +class Outcome: + nh: NativeHome + fake: acp.AcpFake + run: ChatResult + extra: dict[str, Any] = field(default_factory=dict) + + +def _alive(pid: int) -> bool: + try: + with open(f"/proc/{pid}/stat", encoding="utf-8") as fh: + return fh.read().rsplit(")", 1)[1].split()[0] != "Z" + except (FileNotFoundError, ProcessLookupError, IndexError): + return False + + +def _drive(root: Path, row: Row) -> Outcome: + fake = acp.AcpFake(root / "acp", row.turns) + nh = make_home(root, {"provider": "copilot-acp", "default": "copilot-acp"}, env_file=fake.env()) + return Outcome(nh, fake, run_chat(nh, "Answer the question, please.")) + + +@pytest.fixture(scope="module") +def outcomes(tmp_path_factory: pytest.TempPathFactory): + base = tmp_path_factory.mktemp("copilot_acp_errors") + with ThreadPoolExecutor(max_workers=len(ROWS)) as pool: + futures = {name: pool.submit(_drive, base / name, row) for name, row in ROWS.items()} + done = {name: fut.result() for name, fut in futures.items()} + yield done + for out in done.values(): + for pid in out.fake.pids(): + if _alive(pid): + with contextlib.suppress(ProcessLookupError): # exited between the check and the kill + os.kill(pid, signal.SIGKILL) + + +@pytest.mark.parametrize("name", list(ROWS)) +def test_acp_failure_is_retried_per_semantics_and_surfaced_once(outcomes, name): + row, out = ROWS[name], outcomes[name] + run, fake = out.run, out.fake + assert fake.invalid() == [], f"requests rejected by the ACP schema: {fake.invalid()}" + with known_failure(r"^crash_\w+: [3-9] model calls, expected 2", + "#121467 a crash whose stderr lags the exit reads as a timeout and is retried past the budget"): + assert len(fake.main_prompts()) == row.model_calls, ( + f"{name}: {len(fake.main_prompts())} model calls, expected {row.model_calls}\n{run.describe()}") + assert (run.returncode == 0) is row.succeeds, run.describe() + assert run.stdout.count(row.visible) == 1, f"{row.visible!r} must be shown exactly once\n{run.describe()}" + rows = messages(out.nh, latest_session(out.nh)) + persisted = "\n".join(r["content"] or "" for r in rows if r["role"] == "assistant") + for text in row.never_persisted: + assert text not in persisted, f"failed-attempt text {text!r} was persisted: {rows}" + if row.succeeds: + assert persisted.count(OK_TEXT) == 1, rows + pids = fake.pids() + assert len(pids) == row.model_calls, f"one agent process per model call expected: {pids}" + wait_until(lambda: not [p for p in pids if _alive(p)], 10.0, f"agent processes {pids} to exit") + + +REMEDY_RE = re.compile(r"`(hermes [^`]+)`") + + +@pytest.mark.xfail(strict=True, raises=KnownSymptom, reason=KNOWN["auth_remedy"]) +def test_auth_failure_remedy_is_an_actionable_command(outcomes): + """The sign-in remedy printed for an ACP ``Authentication required`` must not be a dead end: any + ``hermes ...`` command it names has to be implemented for this provider.""" + out = outcomes["auth_required_not_retried"] + assert out.run.returncode != 0 and "Authentication required" in out.run.stdout, out.run.describe() + for command in REMEDY_RE.findall(out.run.stdout): + argv = [sys.executable, "-m", "hermes_cli.main", *command.split()[1:]] + proc = subprocess.run(argv, cwd=out.nh.project, env=out.nh.env(), capture_output=True, text=True, + timeout=60, stdin=subprocess.DEVNULL) + said = (proc.stdout + proc.stderr).lower() + if "not implemented" in said: + raise KnownSymptom(f"remedy {command!r} is not implemented for copilot-acp: {said.strip()[:300]}") diff --git a/tests/e2e/core/providers/test_native_copilot_acp_streaming.py b/tests/e2e/core/providers/test_native_copilot_acp_streaming.py new file mode 100644 index 000000000000..0d5894c8f9c7 --- /dev/null +++ b/tests/e2e/core/providers/test_native_copilot_acp_streaming.py @@ -0,0 +1,240 @@ +"""Copilot ACP wire conformance, part 3: a long deep-reasoning turn must stream progress while in flight. + +The fake agent (``tests/fakes/providers/copilot_acp.py``) streams an early message chunk and then +``agent_thought_chunk`` updates for ~4 s before its final chunk and the ``session/prompt`` result, the +shape of a deep-reasoning tier that thinks for minutes. The contract: chunks the agent has already +emitted reach the user's surface BEFORE the turn completes, so a long turn never looks hung. + +Two real surfaces, each timestamped against the fake's own clock (same host): + +* the Desktop/TUI event stream (``python -m tui_gateway.entry`` over stdio): ``reasoning.delta`` / + ``message.delta`` events carrying the agent's text (#120550); +* ACP composition — Hermes itself served as an ACP agent (``hermes acp``) on top of the copilot-acp + provider: the outer client must get ``session/update`` chunks before the inner turn ends (#101507). + +Both are red on main (the ACP client buffers the whole response, then replays it as a stream), so +both are strict xfails that raise :class:`KnownSymptom` only for "no progress before the result". +""" + +from __future__ import annotations + +import json +import os +import queue +import signal +import subprocess +import sys +import threading +import time +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, Callable + +import pytest + +from tests.e2e.core.providers._native_helpers import TURN_TIMEOUT, KnownSymptom, NativeHome, make_home +from tests.fakes.providers import copilot_acp as acp + +pytest.importorskip("acp.schema", reason="the fake validates against the agent-client-protocol package (acp extra)") + +pytestmark = [ + pytest.mark.skipif(not sys.platform.startswith("linux"), reason="POSIX launcher script + /proc pid checks"), + pytest.mark.live_system_guard_bypass, +] + +HEAD, ANSWER, THOUGHT = "DEEP-HEAD ", "DEEP-ANSWER-DONE", "DEEP-THOUGHT" +THOUGHT_STEPS, STEP_S = 8, 0.5 +MIN_SPAN_S = THOUGHT_STEPS * STEP_S * 0.75 # vacuity: the agent really spent seconds before its result +MARKERS = (HEAD.strip(), THOUGHT, ANSWER) + + +KNOWN: dict[str, str] = { + "tui_stream": "#120550 copilot-acp buffers the whole turn: no reasoning/message delta reaches the UI in flight", + "nested_acp": "#101507 hermes acp over copilot-acp forwards inner ACP chunks only after the inner turn ends", +} + + +def _deep_turn() -> list[dict[str, Any]]: + thoughts = [acp.thought(f"{THOUGHT}-{i} ", delay=STEP_S) for i in range(THOUGHT_STEPS)] + return [acp.message(HEAD, delay=0.5), *thoughts, acp.message(ANSWER, delay=STEP_S), acp.result()] + + +@dataclass +class Observed: + fake: acp.AcpFake + received: list[tuple[float, dict[str, Any]]] = field(default_factory=list) # (wall clock, message) + final_text: str = "" + returncode: int | None = None + stderr: str = "" + + +def _pump(proc: subprocess.Popen, sink: "queue.Queue[tuple[float, dict[str, Any]] | None]") -> None: + for line in proc.stdout: # type: ignore[union-attr] + try: + sink.put((time.time(), json.loads(line))) + except json.JSONDecodeError: + continue + sink.put(None) + + +class _LineRpc: + """Newline-delimited JSON-RPC 2.0 over a child's stdio; keeps every inbound message with its arrival time.""" + + def __init__(self, proc: subprocess.Popen, obs: Observed, on_request: Callable[[dict[str, Any]], Any]): + self.proc, self.obs, self.on_request = proc, obs, on_request + self.inbox: queue.Queue[tuple[float, dict[str, Any]] | None] = queue.Queue() + self.next_id = 0 + threading.Thread(target=_pump, args=(proc, self.inbox), daemon=True).start() + + def send(self, msg: dict[str, Any]) -> None: + self.proc.stdin.write(json.dumps(msg) + "\n") # type: ignore[union-attr] + self.proc.stdin.flush() # type: ignore[union-attr] + + def until(self, pred: Callable[[dict[str, Any]], bool], timeout: float = TURN_TIMEOUT) -> dict[str, Any]: + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + try: + item = self.inbox.get(timeout=0.5) + except queue.Empty: + continue + if item is None: + raise AssertionError(f"child closed stdout (rc={self.proc.poll()})") + self.obs.received.append(item) + msg = item[1] + if "method" in msg and "id" in msg: + self.send({"jsonrpc": "2.0", "id": msg["id"], **self.on_request(msg)}) + if pred(msg): + return msg + raise AssertionError(f"timed out after {timeout}s") + + def call(self, method: str, params: dict[str, Any], timeout: float = TURN_TIMEOUT) -> Any: + self.next_id += 1 + rid = self.next_id + self.send({"jsonrpc": "2.0", "id": rid, "method": method, "params": params}) + msg = self.until(lambda m: m.get("id") == rid and "method" not in m, timeout) + assert "error" not in msg, f"{method} failed: {msg['error']}" + return msg.get("result") + + +def _refuse(msg: dict[str, Any]) -> dict[str, Any]: + return {"error": {"code": -32601, "message": f"test client does not implement {msg['method']}"}} + + +def _run_child(argv: list[str], nh: NativeHome, fake: acp.AcpFake, drive: Callable[[_LineRpc], str]) -> Observed: + obs = Observed(fake) + stderr_path = nh.root / "child_stderr.log" + with open(stderr_path, "w", encoding="utf-8") as err: + proc = subprocess.Popen(argv, cwd=nh.project, env=nh.env(), stdin=subprocess.PIPE, stdout=subprocess.PIPE, + stderr=err, text=True, bufsize=1) + try: + obs.final_text = drive(_LineRpc(proc, obs, _refuse)) + finally: + proc.stdin.close() # type: ignore[union-attr] + try: + proc.wait(timeout=60) + except subprocess.TimeoutExpired: + proc.kill() + proc.wait(timeout=10) + obs.returncode, obs.stderr = proc.returncode, stderr_path.read_text(encoding="utf-8")[-3000:] + return obs + + +def _home(root: Path) -> tuple[NativeHome, acp.AcpFake]: + fake = acp.AcpFake(root / "acp", [_deep_turn()]) + return make_home(root, {"provider": "copilot-acp", "default": "copilot-acp"}, env_file=fake.env()), fake + + +def _tui_gateway(root: Path) -> Observed: + nh, fake = _home(root) + + def drive(rpc: _LineRpc) -> str: + def event(etype: str) -> Callable[[dict[str, Any]], bool]: + return lambda m: m.get("method") == "event" and (m.get("params") or {}).get("type") == etype + + rpc.until(event("gateway.ready"), 120) + sid = rpc.call("session.create", {})["session_id"] + rpc.call("prompt.submit", {"session_id": sid, "text": "Think deeply, then answer."}) + done = rpc.until(event("message.complete")) + return str((done["params"].get("payload") or {}).get("text") or "") + + return _run_child([sys.executable, "-m", "tui_gateway.entry"], nh, fake, drive) + + +def _nested_acp(root: Path) -> Observed: + nh, fake = _home(root) + + def drive(rpc: _LineRpc) -> str: + rpc.call("initialize", {"protocolVersion": 1, "clientCapabilities": {}, "clientInfo": {"name": "e2e", "version": "1"}}, 120) + sid = rpc.call("session/new", {"cwd": str(nh.project), "mcpServers": []}, 120)["sessionId"] + result = rpc.call("session/prompt", {"sessionId": sid, "prompt": [{"type": "text", "text": "Think deeply, then answer."}]}) + assert result.get("stopReason") == "end_turn", result + return "".join(_chunk_text(m) for _, m in rpc.obs.received) + + return _run_child([sys.executable, "-m", "hermes_cli.main", "acp"], nh, fake, drive) + + +def _chunk_text(msg: dict[str, Any]) -> str: + params = msg.get("params") or {} + if msg.get("method") == "session/update": # outer ACP surface + content = (params.get("update") or {}).get("content") or {} + return str(content.get("text") or "") if isinstance(content, dict) else "" + if msg.get("method") == "event" and params.get("type") in ("message.delta", "reasoning.delta"): # TUI surface + return str((params.get("payload") or {}).get("text") or "") + return "" + + +SURFACES: dict[str, Callable[[Path], Observed]] = {"tui_stream": _tui_gateway, "nested_acp": _nested_acp} + + +@pytest.fixture(scope="module") +def observed(tmp_path_factory: pytest.TempPathFactory): + base = tmp_path_factory.mktemp("copilot_acp_stream") + with ThreadPoolExecutor(max_workers=len(SURFACES)) as pool: + futures = {name: pool.submit(fn, base / name) for name, fn in SURFACES.items()} + done = {name: fut.result() for name, fut in futures.items()} + yield done + for obs in done.values(): + for pid in obs.fake.pids(): + if os.path.exists(f"/proc/{pid}"): + try: + os.kill(pid, signal.SIGKILL) + except ProcessLookupError: + pass + + +def _agent_timeline(fake: acp.AcpFake) -> tuple[float, float]: + """(first chunk sent, prompt result sent) on the fake agent's clock, for the single main turn.""" + outs = [r for r in fake.records() if r["dir"] == "out" and r["turn"] == 0] + first = min(r["t"] for r in outs if (r["msg"] or {}).get("method") == "session/update") + done = min(r["t"] for r in outs if "stopReason" in ((r["msg"] or {}).get("result") or {})) + return first, done + + +@pytest.mark.parametrize("surface", [ + pytest.param(name, marks=pytest.mark.xfail(strict=True, raises=KnownSymptom, reason=KNOWN[name])) + for name in SURFACES +]) +def test_long_reasoning_turn_streams_progress_before_it_completes(observed, surface): + obs = observed[surface] + assert obs.fake.invalid() == [], f"requests rejected by the ACP schema: {obs.fake.invalid()}" + assert ANSWER in obs.final_text, f"turn did not complete with the agent's answer: {obs.final_text!r}\n{obs.stderr}" + first, done = _agent_timeline(obs.fake) + assert done - first >= MIN_SPAN_S, f"vacuity: the agent streamed for only {done - first:.2f}s" + seen_at = _first_marker_arrival(obs.received) + if seen_at is None or seen_at >= done - 0.3: + late = [round(t - done, 2) for t, m in obs.received if _chunk_text(m)] + raise KnownSymptom(f"{surface}: no agent chunk reached the surface during the {done - first:.1f}s turn; " + f"chunk arrivals relative to the result: {late}") + assert seen_at >= first, "a chunk cannot arrive before the agent sent it" + + +def _first_marker_arrival(received: list[tuple[float, dict[str, Any]]]) -> float | None: + """Arrival time of the delta that first completes one of the agent's markers (surfaces may re-split + chunks, so match on the accumulated text, not per delta).""" + buffer = "" + for t, msg in received: + buffer += _chunk_text(msg) + if any(k in buffer for k in MARKERS): + return t + return None diff --git a/tests/e2e/core/providers/test_native_gemini_compaction.py b/tests/e2e/core/providers/test_native_gemini_compaction.py new file mode 100644 index 000000000000..85ddc6e0b3ca --- /dev/null +++ b/tests/e2e/core/providers/test_native_gemini_compaction.py @@ -0,0 +1,134 @@ +"""Gemini native wire conformance: auto-compaction inside a thought-signature session. + +Turn 1 (process A) reads two bulky files (two signed functionCall steps) and answers. Turn 2 +(process B, ``--resume``) reads a third file; that response reports a huge ``promptTokenCount``, so +Hermes compacts mid-turn (tiny ``compression.threshold_tokens``) before the next step, then the +model makes one more signed call and answers. + +The fake validates every request like Google: roles alternate, every functionResponse follows +its functionCall (count / name / id), and each functionCall step of the current turn carries a +signature Google issued. The tests additionally require that compaction really happened on the +wire and that every functionCall that survives it still carries its own signature verbatim. +""" + +from __future__ import annotations + +import json +import random +from dataclasses import dataclass + +import pytest + +from tests.e2e.core.providers import _native_helpers as nh +from tests.fakes.providers.gemini_native import ( + HERMES_ENV, + Call, + Calls, + GeminiFake, + Recorded, + Text, + hermes_model, +) + +SUMMARY_MARK = "GEMINI-SUMMARY-CHECKPOINT" +ANSWER_1 = "Turn one answer GEMINI-CMP-ONE" +ANSWER_2 = "Turn two answer GEMINI-CMP-TWO" +ECHO = "GEMINI-CMP-ECHO-3391" +BIG_PROMPT = 50_000 # far above threshold_tokens below: compaction must fire + +SUMMARY = (f"## Goal\nKeep helping with the files ({SUMMARY_MARK}).\n## Progress\n### Done\n" + "- Read big1.txt and big2.txt.\n## Next Steps\n- Continue with big3.txt.\n") + + +def _summary_route(rec: Recorded) -> Text | None: + """The compaction summariser call is the one generate call that declares no tools.""" + return None if (rec.body or {}).get("tools") else Text(SUMMARY, signed=False) + + +@dataclass +class Run: + home: nh.NativeHome + fake: GeminiFake + turns: list[nh.ChatResult] + + def calls(self) -> list[Recorded]: + return self.fake.generate_calls() + + def split(self) -> tuple[Recorded, Recorded, list[Recorded]]: + """(last main request before the summary, first after it, every main request after it).""" + calls = self.calls() + idx = next((i for i, c in enumerate(calls) if c.reply.startswith("route:")), None) + assert idx is not None, f"no compaction summary request reached Google: {[c.reply for c in calls]}" + before = [c for c in calls[:idx] if c.reply.startswith("script:")] + after = [c for c in calls[idx + 1:] if c.reply.startswith("script:")] + assert before and after, [c.reply for c in calls] + return before[-1], after[0], after + + +@pytest.fixture(scope="module") +def run(tmp_path_factory: pytest.TempPathFactory) -> Run: + root = tmp_path_factory.mktemp("gemini_compaction") + home = nh.make_home(root, hermes_model(context_length=64_000), env_file=HERMES_ENV, + extra_config={"compression": {"threshold_tokens": 12_000, "protect_last_n": 4}}) + rng = random.Random(7) + words = "alpha bravo charlie delta echo foxtrot golf hotel india juliet kilo lima mike oscar".split() + paths = [] + for i in range(1, 4): + path = home.project / f"big{i}.txt" + lines = [" ".join(rng.choice(words) for _ in range(14)) for _ in range(260)] + path.write_text("\n".join(lines) + f"\nBIG-END-{i}\n", encoding="utf-8") + paths.append(str(path)) + script = [ + Calls([Call("read_file", {"path": paths[0]})]), + Calls([Call("read_file", {"path": paths[1]})]), + Text(ANSWER_1), + Calls([Call("read_file", {"path": paths[2]})], prompt_tokens=BIG_PROMPT), + Calls([Call("terminal", {"command": f"echo {ECHO}"})], prompt_tokens=BIG_PROMPT), + Text(ANSWER_2, prompt_tokens=3_000), + ] + with GeminiFake(root / "fake", script, route=_summary_route) as fake: + # Outcomes are asserted by the tests (a rejected request must name the broken contract). + first = nh.run_chat(home, "Read big1.txt and big2.txt.", env=fake.child_env()) + second = nh.run_chat(home, "Now read big3.txt, then echo the marker.", env=fake.child_env(), + resume=nh.latest_session(home)) + return Run(home, fake, [first, second]) + + +def test_compaction_happens_on_the_wire(run: Run) -> None: + before, after, _ = run.split() + # Without compaction the next request = previous request + (functionCall, functionResponse). + assert len(after.contents) < len(before.contents) + 2, (len(before.contents), len(after.contents)) + assert SUMMARY_MARK in after.all_text(), "compaction summary never reached the next request" + assert all(t.returncode == 0 for t in run.turns), [t.describe() for t in run.turns] + assert ANSWER_2 in run.turns[1].stdout, run.turns[1].describe() + + +def test_post_compaction_requests_are_valid_and_signed(run: Run) -> None: + """No request (before or after compaction) was rejected, and every functionCall still on the + wire after compaction carries the exact signature Google minted for that call id.""" + assert run.fake.rejections() == [], run.fake.rejections() + _, _, after = run.split() + issued = run.fake.call_signatures + for rec in after: + call_ids = {p["functionCall"].get("id") for p in rec.parts("functionCall")} + for part in rec.parts("functionCall"): + cid = part["functionCall"].get("id") + assert cid in issued, f"functionCall id {cid!r} was never issued by Google" + assert part.get("thoughtSignature") == issued[cid], f"signature for {cid} lost/altered: {part}" + for part in rec.parts("functionResponse"): + assert part["functionResponse"].get("id") in call_ids, f"orphan functionResponse: {part}" + last = after[-1] + echoed = [p["functionResponse"] for p in last.parts("functionResponse")] + assert any(ECHO in json.dumps(r["response"]) for r in echoed), echoed + + +def test_persisted_transcript_keeps_pairs_and_signatures(run: Run) -> None: + rows = nh.messages(run.home) # active rows across the (possibly split) session lineage + open_ids: set[str] = set() + for row in rows: + for tc in nh.tool_calls_of(row) if row["role"] == "assistant" else []: + open_ids.add(str(tc.get("id"))) + if row["role"] == "tool": + assert row.get("tool_call_id") in open_ids, f"orphan tool row {row.get('tool_call_id')}" + nh.assert_no_duplicate_assistant_text(rows, ANSWER_2) + assert any(ANSWER_2 in (r.get("content") or "") for r in rows if r["role"] == "assistant") diff --git a/tests/e2e/core/providers/test_native_gemini_errors.py b/tests/e2e/core/providers/test_native_gemini_errors.py new file mode 100644 index 000000000000..94807db1f814 --- /dev/null +++ b/tests/e2e/core/providers/test_native_gemini_errors.py @@ -0,0 +1,159 @@ +"""Gemini native wire conformance: documented errors, blocked candidates and a mid-stream drop. + +Each scenario is one real ``hermes chat -q`` turn against the fake Google endpoint +(``tests/fakes/providers/gemini_native.py``); the scenarios are independent and run concurrently. + +* retryable (429 ``RESOURCE_EXHAUSTED`` with ``RetryInfo``, 503 ``UNAVAILABLE``): retried, then the + answer is printed and persisted once; +* terminal (400 ``INVALID_ARGUMENT``, candidate ``finishReason`` ``SAFETY`` / ``RECITATION``, a prompt + blocked through ``promptFeedback.blockReason``): exactly one request, surfaced once, no fake + success persisted; +* a TLS connection dropped mid-SSE: recovered with no partial/duplicated assistant row. +""" + +from __future__ import annotations + +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path + +import pytest + +from tests.e2e.core.providers import _native_helpers as nh +from tests.fakes.providers.gemini_native import ( + HERMES_ENV, + Blocked, + Drop, + GeminiFake, + GoogleError, + Recorded, + Reply, + Responder, + Text, + hermes_model, +) + +KNOWN = { + "prompt_blocked": "#121317 promptFeedback.blockReason is retried 9x as an empty stream and reported as " + "'temporarily unavailable'", +} + +NEVER = "GEMINI-MUST-NOT-BE-SHOWN" +QUOTA = "Resource has been exhausted (e.g. check quota)." +PARTIAL = "GEMINI-PARTIAL-DROPPED " + + +@dataclass +class Case: + script: list[Reply] + route: Responder | None = None + max_retries: int = 2 + + +CASES: dict[str, Case] = { + "rate_limited_429": Case([GoogleError(429, "RESOURCE_EXHAUSTED", QUOTA, 1), + GoogleError(429, "RESOURCE_EXHAUSTED", QUOTA, 1), + Text("GEMINI-RECOVERED-429")], max_retries=3), + "unavailable_503": Case([GoogleError(503, "UNAVAILABLE", "The model is overloaded. Please try again later."), + Text("GEMINI-RECOVERED-503")]), + "invalid_argument_400": Case([GoogleError(400, "INVALID_ARGUMENT", "Request contains an invalid argument. " + "GEMINI-MARK-400"), Text(NEVER)]), + "safety": Case([Blocked("SAFETY"), Text(NEVER)]), + "recitation": Case([Blocked("RECITATION"), Text(NEVER)]), + # Google blocks the same prompt every time: answer every attempt with the block. + "prompt_blocked": Case([], route=lambda rec: Blocked("SAFETY", prompt=True), max_retries=3), + "stream_drop": Case([Drop(PARTIAL), Text("GEMINI-FULL-AFTER-DROP")]), +} + + +@dataclass +class Outcome: + result: nh.ChatResult + calls: list[Recorded] + rows: list[dict] + + @property + def output(self) -> str: + return self.result.stdout + self.result.stderr + + def describe(self) -> str: + return f"{self.result.describe()}\ncalls={[(c.reply, c.status) for c in self.calls]}" + + +def _run(root: Path, name: str, case: Case) -> Outcome: + home = nh.make_home(root / name, hermes_model(), env_file=HERMES_ENV, + extra_config={"agent": {"api_max_retries": case.max_retries}}) + with GeminiFake(root / name / "fake", case.script, route=case.route) as fake: + result = nh.run_chat(home, f"Say hello ({name}).", env=fake.child_env()) + calls = fake.generate_calls() + return Outcome(result, calls, nh.messages(home)) + + +@pytest.fixture(scope="module") +def outcomes(tmp_path_factory: pytest.TempPathFactory) -> dict[str, Outcome]: + base = tmp_path_factory.mktemp("gemini_errors") + with ThreadPoolExecutor(len(CASES)) as pool: + futures = {name: pool.submit(_run, base, name, case) for name, case in CASES.items()} + return {name: f.result() for name, f in futures.items()} + + +def _answer_once(o: Outcome, answer: str) -> None: + assert o.result.returncode == 0, o.describe() + assert o.result.stdout.count(answer) == 1, o.describe() + assistant = [r for r in o.rows if r["role"] == "assistant" and r.get("content")] + assert [r["content"] for r in assistant] == [answer], assistant + + +RETRYABLE = { # case -> (statuses the fake must have served, in order; the answer) + "rate_limited_429": ([429, 429, 200], "GEMINI-RECOVERED-429"), + "unavailable_503": ([503, 200], "GEMINI-RECOVERED-503"), +} + + +@pytest.mark.parametrize("name", list(RETRYABLE)) +def test_retryable_error_is_retried_then_succeeds(outcomes: dict[str, Outcome], name: str) -> None: + o = outcomes[name] + statuses, answer = RETRYABLE[name] + assert [c.status for c in o.calls] == statuses, o.describe() + _answer_once(o, answer) + + +TERMINAL = { # case -> a word the surfaced error must carry (None: any visible error) + "invalid_argument_400": "GEMINI-MARK-400", + "safety": "safety", + "recitation": None, + "prompt_blocked": "block", +} + + +def _terminal_param(name: str): + marks = [pytest.mark.xfail(strict=True, raises=nh.KnownSymptom, reason=KNOWN[name])] if name in KNOWN else [] + return pytest.param(name, marks=marks, id=name) + + +@pytest.mark.parametrize("name", [_terminal_param(n) for n in TERMINAL]) +def test_non_retryable_surfaced_once(outcomes: dict[str, Outcome], name: str) -> None: + o = outcomes[name] + assert o.calls, f"no request reached the fake\n{o.describe()}" + if len(o.calls) != 1: + # The tracked bug's symptom for KNOWN cases; a plain failure for every other terminal case. + failure = nh.KnownSymptom if name in KNOWN else AssertionError + raise failure(f"terminal Google response was re-sent {len(o.calls)}x\n{o.describe()}") + assert o.result.returncode != 0, o.describe() + word = TERMINAL[name] + lines = [ln.strip() for ln in o.result.stdout.splitlines() if ln.strip()] + assert lines, f"nothing surfaced to the user\n{o.describe()}" + assert len(lines) == len(set(lines)), f"error surfaced more than once\n{o.describe()}" + if word: + assert word.lower() in o.output.lower(), o.describe() + assert NEVER not in o.output + assert not [r for r in o.rows if r["role"] == "assistant" and NEVER in (r.get("content") or "")] + + +def test_stream_drop_recovers_without_duplicate_rows(outcomes: dict[str, Outcome]) -> None: + o = outcomes["stream_drop"] + assert [c.reply for c in o.calls] == ["script:Drop", "script:Text"], o.describe() + assert all(c.stream for c in o.calls) + _answer_once(o, "GEMINI-FULL-AFTER-DROP") + assert not [r for r in o.rows if PARTIAL.strip() in (r.get("content") or "")], o.rows + nh.assert_no_duplicate_assistant_text(o.rows, "GEMINI-FULL-AFTER-DROP") diff --git a/tests/e2e/core/providers/test_native_gemini_schema.py b/tests/e2e/core/providers/test_native_gemini_schema.py new file mode 100644 index 000000000000..1fec461b2b04 --- /dev/null +++ b/tests/e2e/core/providers/test_native_gemini_schema.py @@ -0,0 +1,187 @@ +"""Gemini native wire conformance: MCP tool schemas Google's parser rejects arrive sanitized. + +An MCP server (a tiny stdio JSON-RPC process) exposes one tool whose ``inputSchema`` uses JSON +Schema keywords that Google's ``FunctionDeclaration`` parser does not accept (``$schema``, +``$ref``/``$defs``, ``additionalProperties``, ``oneOf``, ``const``, list-valued ``type``, integer +``enum``, a ``required`` entry naming no property). A real ``hermes chat -q`` turn declares it to +the fake Google endpoint, which rejects anything Google would (HTTP 400 INVALID_ARGUMENT) and +then drives a call to the tool so the round trip is proven end to end. + +Two API surfaces: ``v1beta`` (default; ``parametersJsonSchema``) and ``v1`` (pinned through +``model.base_url``; the proto ``Schema`` subset in ``parameters``, validated field by field). +""" + +from __future__ import annotations + +import json +import sys +import textwrap +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import pytest + +from tests.e2e.core.providers import _native_helpers as nh +from tests.fakes.providers.gemini_native import HERMES_ENV, Call, Calls, GeminiFake, Recorded, Text, hermes_model + +KNOWN = { + "ref_dropped_v1": "#99438 legacy `parameters` path drops $ref/$defs instead of inlining (empty schema)", + "array_items_v1": "#71804 array parameter without `items` is sent as-is; Google 400s 'items: missing field'", +} + +TOOL = "mcp__hostile__lookup" +MCP_CANARY = "GEMINI-MCP-CANARY-7731" +DONE = "GEMINI-MCP-DONE" +V1 = "https://generativelanguage.googleapis.com/v1" +_FORBIDDEN_ANYWHERE = ("$ref", "$schema", "$defs") + +HOSTILE: dict[str, Any] = { + "$schema": "http://json-schema.org/draft-07/schema#", + "type": "object", + "additionalProperties": False, + "$defs": {"Mode": {"type": "string", "enum": ["fast", "slow"]}}, + "properties": { + "query": {"type": "string", "minLength": 1, "examples": ["x"]}, + "mode": {"$ref": "#/$defs/Mode"}, + "limit": {"type": ["integer", "null"], "exclusiveMinimum": 0, "enum": [1, 5, 10]}, + "filters": {"type": "array", "items": {"oneOf": [ + {"type": "string"}, + {"type": "object", "properties": {"k": {"type": "string"}}, "additionalProperties": False}]}}, + "flag": {"const": True}, + }, + "required": ["query", "not_a_property"], +} +ITEMLESS: dict[str, Any] = {"type": "object", "properties": {"tags": {"type": "array", "description": "tags"}}} + +_MCP_SERVER = textwrap.dedent(''' + import json, sys + TOOLS = json.loads(sys.argv[1]) + CANARY = sys.argv[2] + for line in sys.stdin: + msg = json.loads(line) + mid, method = msg.get("id"), msg.get("method") + if mid is None: + continue + if method == "initialize": + res = {"protocolVersion": msg["params"].get("protocolVersion", "2025-06-18"), + "capabilities": {"tools": {}}, "serverInfo": {"name": "hostile", "version": "1"}} + elif method == "tools/list": + res = {"tools": TOOLS} + elif method == "tools/call": + args = msg["params"].get("arguments") or {} + res = {"content": [{"type": "text", "text": CANARY + " query=" + str(args.get("query"))}]} + elif method == "ping": + res = {} + else: + err = {"code": -32601, "message": "method not found: " + str(method)} + print(json.dumps({"jsonrpc": "2.0", "id": mid, "error": err}), flush=True) + continue + print(json.dumps({"jsonrpc": "2.0", "id": mid, "result": res}), flush=True) +''') + + +@dataclass +class Outcome: + result: nh.ChatResult + calls: list[Recorded] + rejections: list[str] + + def declaration(self) -> dict[str, Any]: + assert self.calls, self.result.describe() + decl = self.calls[0].declarations().get(TOOL) + assert decl is not None, f"{TOOL} not declared: {sorted(self.calls[0].declarations())}" + return decl + + def function_responses(self) -> list[dict[str, Any]]: + return [p["functionResponse"] for rec in self.calls for p in rec.parts("functionResponse")] + + +def _scenario(root: Path, schema: dict[str, Any], base_url: str | None) -> Outcome: + root.mkdir(parents=True) + server = root / "hostile_mcp.py" + server.write_text(_MCP_SERVER, encoding="utf-8") + tools = [{"name": "lookup", "description": "Look something up.", "inputSchema": schema}] + extra = { + # Keep the MCP tool in the model-facing array (not behind the tool_search bridge) and let + # discovery finish before the one-shot turn is built. + "tools": {"tool_search": {"enabled": "off"}}, + "mcp_discovery_timeout": 60, + "mcp_single_query_discovery_timeout": 60, + "mcp_servers": {"hostile": {"command": sys.executable, + "args": [str(server), json.dumps(tools), MCP_CANARY]}}, + } + home = nh.make_home(root, hermes_model(base_url), env_file=HERMES_ENV, extra_config=extra) + script = [Calls([Call(TOOL, {"query": "q1", "mode": "fast"})]), Text(DONE)] + with GeminiFake(root / "fake", script) as fake: + result = nh.run_chat(home, "Use the lookup tool for q1.", env=fake.child_env()) + return Outcome(result, fake.generate_calls(), fake.rejections()) + + +SCENARIOS = {"v1beta": (HOSTILE, None), "v1": (HOSTILE, V1), "v1_itemless": (ITEMLESS, V1)} + + +@pytest.fixture(scope="module") +def outcomes(tmp_path_factory: pytest.TempPathFactory) -> dict[str, Outcome]: + base = tmp_path_factory.mktemp("gemini_schema") + with ThreadPoolExecutor(len(SCENARIOS)) as pool: + futures = {k: pool.submit(_scenario, base / k, *v) for k, v in SCENARIOS.items()} + return {k: f.result() for k, f in futures.items()} + + +def _assert_round_trip(o: Outcome, version: str) -> None: + assert o.rejections == [], o.rejections + assert [c.version for c in o.calls] == [version, version], [(c.version, c.status) for c in o.calls] + responses = [r for r in o.function_responses() if r.get("name") == TOOL] + assert responses and MCP_CANARY in json.dumps(responses[0]["response"]), o.function_responses() + assert f"{MCP_CANARY} query=q1" in json.dumps(responses[0]["response"]), responses[0] + assert o.result.returncode == 0 and DONE in o.result.stdout, o.result.describe() + + +def _walk(node: Any): + if isinstance(node, dict): + yield node + for v in node.values(): + yield from _walk(v) + elif isinstance(node, list): + for v in node: + yield from _walk(v) + + +def test_v1beta_json_schema_declaration_is_sanitized(outcomes: dict[str, Outcome]) -> None: + o = outcomes["v1beta"] + _assert_round_trip(o, "v1beta") + schema = o.declaration().get("parametersJsonSchema") + assert isinstance(schema, dict), o.declaration() + leaked = [k for node in _walk(schema) for k in node if k in _FORBIDDEN_ANYWHERE] + assert leaked == [], f"reference/meta keywords reached Google: {leaked} in {schema}" + assert schema["properties"]["mode"].get("enum") == ["fast", "slow"], schema["properties"]["mode"] + assert set(schema.get("required") or []) <= set(schema["properties"]), schema.get("required") + + +def test_v1_proto_schema_declaration_is_accepted(outcomes: dict[str, Outcome]) -> None: + """Every field of ``parameters`` parses as Google's proto ``Schema`` (the fake 400s otherwise).""" + o = outcomes["v1"] + _assert_round_trip(o, "v1") + params = o.declaration().get("parameters") + assert isinstance(params, dict) and "parametersJsonSchema" not in o.declaration(), o.declaration() + assert params["required"] == ["query"], params + + +@pytest.mark.xfail(strict=True, raises=nh.KnownSymptom, reason=KNOWN["ref_dropped_v1"]) +def test_v1_ref_parameter_keeps_its_shape(outcomes: dict[str, Outcome]) -> None: + _assert_round_trip(outcomes["v1"], "v1") + mode = outcomes["v1"].declaration()["parameters"]["properties"]["mode"] + if mode == {}: + raise nh.KnownSymptom(f"$ref-typed parameter lost its shape on the v1 wire: {mode}") + assert mode.get("enum") == ["fast", "slow"], mode + + +@pytest.mark.xfail(strict=True, raises=nh.KnownSymptom, reason=KNOWN["array_items_v1"]) +def test_v1_array_without_items_is_accepted(outcomes: dict[str, Outcome]) -> None: + o = outcomes["v1_itemless"] + assert o.calls, o.result.describe() + if any("items: missing field" in r for r in o.rejections): + raise nh.KnownSymptom(f"Google rejected the item-less array: {o.rejections}") + _assert_round_trip(o, "v1") diff --git a/tests/e2e/core/providers/test_native_gemini_tools.py b/tests/e2e/core/providers/test_native_gemini_tools.py new file mode 100644 index 000000000000..e131a4757066 --- /dev/null +++ b/tests/e2e/core/providers/test_native_gemini_tools.py @@ -0,0 +1,133 @@ +"""Gemini native wire conformance: the tool loop and thought-signature replay across ``--resume``. + +Real ``hermes chat -q`` subprocesses talk to Google AI Studio's native ``streamGenerateContent`` +dialect; only the vendor is faked (``tests/fakes/providers/gemini_native.py``, a TLS-terminating +loopback proxy that validates every request like Google and rejects it with a 400 when it would). + +Turn 1 (process A): the model calls ``read_file`` on a seeded file, then answers. +Turn 2 (process B, ``--resume``): the model calls ``terminal``, then answers. +Gemini 3 needs the ``thoughtSignature`` of every functionCall step sent back verbatim; the fake +enforces it for the current turn and the tests assert it for the resumed (previous-turn) history. +""" + +from __future__ import annotations + +import json +from dataclasses import dataclass + +import pytest + +from tests.e2e.core.providers import _native_helpers as nh +from tests.fakes.providers.gemini_native import ( + API_KEY, + HERMES_ENV, + MODEL_ID, + Call, + Calls, + GeminiFake, + Recorded, + Text, + hermes_model, +) + +SEED = "GEMINI-SEED-CANARY-5521" +ECHO = "GEMINI-RESUME-ECHO-8813" +ANSWER_1 = "Turn one answer GEMINI-ANS-ONE" +ANSWER_2 = "Turn two answer GEMINI-ANS-TWO" + + +@dataclass +class Run: + home: nh.NativeHome + fake: GeminiFake + turns: list[nh.ChatResult] + session_id: str + + def main(self) -> list[Recorded]: + return self.fake.main_calls() + + +@pytest.fixture(scope="module") +def run(tmp_path_factory: pytest.TempPathFactory) -> Run: + root = tmp_path_factory.mktemp("gemini_tools") + home = nh.make_home(root, hermes_model(), env_file=HERMES_ENV) + seed = home.project / "seed.txt" + seed.write_text(f"{SEED}\n", encoding="utf-8") + script = [ + Calls([Call("read_file", {"path": str(seed)})]), + Text(ANSWER_1, thought="Summarising the file."), + Calls([Call("terminal", {"command": f"echo {ECHO}"})]), + Text(ANSWER_2), + ] + with GeminiFake(root / "fake", script) as fake: + # No asserts here: a failing turn must show up as the specific test assertion (a rejected + # request, a missing signature), not as a fixture error that hides which contract broke. + first = nh.run_chat(home, "Read seed.txt and tell me what it says.", env=fake.child_env()) + sid = nh.latest_session(home) + second = nh.run_chat(home, "Now echo the marker in the terminal.", env=fake.child_env(), resume=sid) + return Run(home, fake, [first, second], sid) + + +def _call_parts(rec: Recorded) -> dict[str, dict]: + """functionCall id -> the whole Part (so the sibling ``thoughtSignature`` is visible).""" + return {p["functionCall"].get("id"): p for p in rec.parts("functionCall")} + + +def _responses(rec: Recorded) -> dict[str, dict]: + return {p["functionResponse"].get("id"): p["functionResponse"] for p in rec.parts("functionResponse")} + + +def test_function_call_round_trip_pairs_response_and_persists(run: Run) -> None: + """a. functionCall -> real tool -> functionResponse (same name + id) -> final answer printed.""" + assert run.fake.rejections() == [], run.fake.rejections() + main = run.main() + assert len(main) == 4, [(r.reply, r.status) for r in main] + assert all(r.stream and r.query.get("alt") == ["sse"] for r in main) + (call_id,) = list(run.fake.call_signatures)[:1] + follow_up = main[1] + resp = _responses(follow_up).get(call_id) + assert resp is not None, f"no functionResponse for {call_id}: {follow_up.parts('functionResponse')}" + assert resp["name"] == "read_file" + assert SEED in json.dumps(resp["response"]), resp + part = _call_parts(follow_up)[call_id] + assert part.get("thoughtSignature") == run.fake.call_signatures[call_id], part + assert run.turns[0].returncode == 0 and ANSWER_1 in run.turns[0].stdout, run.turns[0].describe() + + rows = nh.messages(run.home, run.session_id) + ids = [tc.get("id") for r in rows if r["role"] == "assistant" for tc in nh.tool_calls_of(r)] + assert call_id in ids, ids + tool_rows = [r for r in rows if r["role"] == "tool" and r.get("tool_call_id") == call_id] + assert len(tool_rows) == 1 and SEED in (tool_rows[0].get("content") or ""), tool_rows + nh.assert_no_duplicate_assistant_text(rows, ANSWER_1) + + +def test_resume_replays_thought_signatures_verbatim(run: Run) -> None: + """b. After ``--resume`` in a new process, turn 1's functionCall goes back with the exact + signature Google issued (not dropped, not a skip-validator dummy), still paired to its result; + the resumed turn's own call is signed too (the fake 400s a missing current-turn signature).""" + assert run.fake.rejections() == [], run.fake.rejections() + first_id, second_id = list(run.fake.call_signatures) + resumed = run.main()[2] + replayed = _call_parts(resumed).get(first_id) + assert replayed is not None, f"turn-1 functionCall {first_id} missing after resume: {resumed.contents}" + assert replayed.get("thoughtSignature") == run.fake.call_signatures[first_id], replayed + assert SEED in json.dumps(_responses(resumed).get(first_id)), resumed.parts("functionResponse") + assert ANSWER_1 in resumed.all_text() + + final = run.main()[3] + assert _call_parts(final)[second_id].get("thoughtSignature") == run.fake.call_signatures[second_id] + assert ECHO in json.dumps(_responses(final).get(second_id)) + assert run.turns[1].returncode == 0 and ANSWER_2 in run.turns[1].stdout, run.turns[1].describe() + nh.assert_no_duplicate_assistant_text(nh.messages(run.home), ANSWER_2) + + +def test_every_request_authenticates_and_targets_the_configured_model(run: Run) -> None: + """The key from ``.env`` travels as ``x-goog-api-key`` (or ``key=``) on every generate call, to + ``/v1beta/models/``; nothing else was requested from the Google host.""" + calls = run.fake.generate_calls() + assert calls + for rec in calls: + key = rec.headers.get("x-goog-api-key") or (rec.query.get("key") or [""])[0] + assert key == API_KEY, f"generate call without the configured API key: {rec.headers}" + assert (rec.version, rec.model) == ("v1beta", MODEL_ID), rec.path + assert [r.path for r in run.fake.requests if not r.rpc] == [] diff --git a/tests/e2e/core/providers/test_native_vertex_errors.py b/tests/e2e/core/providers/test_native_vertex_errors.py new file mode 100644 index 000000000000..5f8087458fe1 --- /dev/null +++ b/tests/e2e/core/providers/test_native_vertex_errors.py @@ -0,0 +1,153 @@ +"""Vertex AI documented errors through the real CLI: retried per semantics or surfaced exactly once. + +Each scenario is one ``hermes chat -q`` turn against its own fake (``tests/fakes/providers/vertex.py``); +all run concurrently. Vertex's OpenAI-compatible endpoint returns google.rpc errors in a list-wrapped +envelope (``[{"error": {"code", "message", "status"}}]``); the fake uses exactly that. + +* 429 RESOURCE_EXHAUSTED (quota) and 503 UNAVAILABLE are transient: retried, then the answer. +* 400 INVALID_ARGUMENT and 403 PERMISSION_DENIED are terminal: one request, message shown once. +* 401 UNAUTHENTICATED on every bearer: one token refresh + retry at most, then shown once. +* OAuth ``invalid_grant`` at the token endpoint: nothing is sent to Vertex; actionable error. +""" + +from __future__ import annotations + +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any + +import pytest + +pytest.importorskip("google.auth", reason="Vertex minting needs google-auth (CI installs it)") + +from tests.e2e.core.providers._native_helpers import ChatResult, KnownSymptom, make_home, run_chat # noqa: E402 +from tests.fakes.providers.vertex import ( # noqa: E402 + PROJECT, + REGION, + SA_EMBEDDED_PROJECT, + Fail, + FakeVertex, + Say, + TokenPolicy, + hermes_setup, +) + +KNOWN: dict[str, str] = { + "guidance:permission_denied": "#121295 Vertex 401/403 are reported as 'rejected your API key' (Vertex has no API " + "key; the fix is the service account / IAM role)", + "guidance:unauthenticated": "#121295 Vertex 401/403 are reported as 'rejected your API key' (Vertex has no API " + "key; the fix is the service account / IAM role)", +} + +QUOTA_MSG = ("Resource exhausted. Please try again later. Please refer to " + "https://cloud.google.com/vertex-ai/generative-ai/docs/error-code-429 for more details.") +UNAVAILABLE_MSG = "The service is currently unavailable." +INVALID_MSG = "Request contains an invalid argument." +DENIED_MSG = (f"Permission 'aiplatform.endpoints.predict' denied on resource '//aiplatform.googleapis.com/projects/" + f"{PROJECT}/locations/{REGION}/publishers/google/models/gemini-3-flash-preview' (or it may not exist).") +UNAUTH_FRAGMENT = "Request had invalid authentication credentials" + + +@dataclass +class Case: + script: list[Any] + policy: TokenPolicy = field(default_factory=TokenPolicy) + answer: str | None = None # the final answer when the error is transient + max_attempts: int = 2 # agent.api_max_retries (total attempts per API call) + + +CASES: dict[str, Case] = { + "rate_limited": Case([Fail(429, QUOTA_MSG), Fail(429, QUOTA_MSG), Say("Answer after the quota recovered.")], + answer="Answer after the quota recovered.", max_attempts=3), + "unavailable": Case([Fail(503, UNAVAILABLE_MSG), Say("Answer after UNAVAILABLE.")], answer="Answer after UNAVAILABLE."), + "invalid_argument": Case([Fail(400, INVALID_MSG)]), + "permission_denied": Case([Fail(403, DENIED_MSG)]), + "unauthenticated": Case([Say("never served")], TokenPolicy(reject_bearers=True)), + "invalid_grant": Case([Say("never served")], TokenPolicy(error=(400, "invalid_grant", "Invalid JWT Signature."))), +} + + +def _run(tmp: Path, name: str, case: Case) -> dict[str, Any]: + fake = FakeVertex(tmp / name / "fake", project=PROJECT, region=REGION, sa_project=SA_EMBEDDED_PROJECT, + script=list(case.script)) + fake.start() + fake.token_policy = case.policy + nh = make_home(tmp / name / "h", **hermes_setup(fake, extra_config={"agent": {"api_max_retries": case.max_attempts}})) + return {"fake": fake, "nh": nh, "turn": run_chat(nh, "Say hello.", env=fake.child_env(), args=("-t", "file"))} + + +@pytest.fixture(scope="module") +def results(tmp_path_factory: pytest.TempPathFactory) -> Any: + tmp = tmp_path_factory.mktemp("vertex_errors") + with ThreadPoolExecutor(max_workers=len(CASES)) as pool: + futures = {name: pool.submit(_run, tmp, name, case) for name, case in CASES.items()} + out = {name: f.result() for name, f in futures.items()} + yield out + for res in out.values(): + res["fake"].stop() + + +def _output(turn: ChatResult) -> str: + return turn.stdout + turn.stderr + + +@pytest.mark.parametrize("name", ["rate_limited", "unavailable"]) +def test_transient_error_is_retried_then_answers(results: dict[str, Any], name: str) -> None: + res, case = results[name], CASES[name] + fake, turn = res["fake"], res["turn"] + assert turn.returncode == 0 and case.answer in turn.stdout, turn.describe() + statuses = [r.get("status") for r in fake.requests] + assert statuses == [f.status for f in case.script if isinstance(f, Fail)] + [200], statuses + bodies = [r["body"]["messages"] for r in fake.requests] + assert all(b == bodies[0] for b in bodies), "a retry changed the conversation it resent" + assert len({r["auth"] for r in fake.requests}) == 1 and len(fake.token_requests) == 1, "retry re-minted needlessly" + + +@pytest.mark.parametrize(("name", "vendor_text"), [("invalid_argument", INVALID_MSG), ("permission_denied", DENIED_MSG)]) +def test_terminal_error_surfaced_once_without_retry(results: dict[str, Any], name: str, vendor_text: str) -> None: + res = results[name] + fake, turn = res["fake"], res["turn"] + assert len(fake.requests) == 1, f"terminal {name} was retried: {[r.get('status') for r in fake.requests]}" + assert turn.returncode != 0, turn.describe() + assert _output(turn).count(vendor_text) == 1, f"vendor message not surfaced exactly once:\n{turn.describe()}" + + +def test_rejected_bearer_refreshes_once_then_surfaces(results: dict[str, Any]) -> None: + """Every bearer 401s: Hermes may re-mint and retry once, never loop, and shows the error once.""" + res = results["unauthenticated"] + fake, turn = res["fake"], res["turn"] + statuses = [r.get("status") for r in fake.requests] + assert statuses and set(statuses) == {401} and len(statuses) <= 2, statuses + assert turn.returncode != 0 + assert _output(turn).count(UNAUTH_FRAGMENT) == 1, turn.describe() + + +def test_oauth_invalid_grant_sends_nothing_and_names_the_credential(results: dict[str, Any]) -> None: + """The SA key is refused at Google's token endpoint: no Vertex request goes out with a missing + bearer, and the user is told which credential setting to fix.""" + res = results["invalid_grant"] + fake, turn = res["fake"], res["turn"] + assert fake.token_requests and all(t.get("error") == "invalid_grant" for t in fake.token_requests) + assert fake.token_requests[0]["claims"], "the JWT assertion itself failed verification" + assert not fake.requests, f"requests reached Vertex without a token: {[r['auth'][:16] for r in fake.requests]}" + assert turn.returncode != 0 + out = _output(turn) + assert "VERTEX_CREDENTIALS_PATH" in out or "GOOGLE_APPLICATION_CREDENTIALS" in out, turn.describe() + + +def _guidance_param(name: str) -> Any: + return pytest.param(name, marks=pytest.mark.xfail(strict=True, raises=KnownSymptom, reason=KNOWN[f"guidance:{name}"])) + + +@pytest.mark.parametrize("name", [_guidance_param("permission_denied"), _guidance_param("unauthenticated")]) +def test_auth_failure_guidance_is_vertex_specific(results: dict[str, Any], name: str) -> None: + """Vertex authenticates with an OAuth service account, so the guidance must not send the user to + rotate an 'API key' that does not exist.""" + res = results[name] + turn = res["turn"] + if not (res["fake"].requests and turn.returncode != 0): + raise RuntimeError(f"{name}: the auth failure never happened:\n{turn.describe()}") + guidance = _output(turn).split("Provider said:")[0].lower() + if "api key" in guidance: + raise KnownSymptom(f"Vertex auth failure blamed on an API key:\n{turn.stdout}") diff --git a/tests/e2e/core/providers/test_native_vertex_recovery.py b/tests/e2e/core/providers/test_native_vertex_recovery.py new file mode 100644 index 000000000000..7eff31dfc958 --- /dev/null +++ b/tests/e2e/core/providers/test_native_vertex_recovery.py @@ -0,0 +1,169 @@ +"""Vertex AI recovery paths: compaction in a signed reasoning session, and a mid-stream drop. + +Real ``hermes chat -q`` turns against the fake Vertex (``tests/fakes/providers/vertex.py``), which +rejects exactly what Gemini 3 rejects after history surgery: orphaned tool results, unanswered +calls, and a current-turn function call without its thought signature (or with a signature it +never issued). Scenarios run concurrently in a module fixture: + +* ``compaction`` — ``--reasoning high``, every step a signed ``read_file`` call, a tiny + ``compression.threshold_tokens`` and large reported prompt tokens, across a ``--resume``. +* ``drop`` — the stream after a signed tool step is cut mid-body (incomplete chunked + TLS response); the retry must resend a valid request and persist the answer once. +""" + +from __future__ import annotations + +import json +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from typing import Any, Callable + +import pytest + +pytest.importorskip("google.auth", reason="Vertex minting needs google-auth (CI installs it)") + +from tests.e2e.core.providers._native_helpers import ( # noqa: E402 + ChatResult, + assert_no_duplicate_assistant_text, + latest_session, + make_home, + messages, + run_chat, + tool_calls_of, +) +from tests.fakes.providers.vertex import ( # noqa: E402 + PROJECT, + REGION, + SA_EMBEDDED_PROJECT, + Call, + Drop, + FakeVertex, + Say, + hermes_setup, +) + +KNOWN: dict[str, str] = {} + +SUMMARY_MARK = "VERTEX-SUMMARY-OK" +TURN1_FINAL = "Compaction turn one done." +TURN2_FINAL = "Compaction turn two done." +DROP_PARTIAL = "DROPPED-PARTIAL-ANSWER that never finishes" +DROP_FINAL = "Recovered answer after the stream dropped." +# Reported prompt tokens stay far above ``threshold_tokens`` so the compressor must run. +CONTEXT_LENGTH = 64_000 +THRESHOLD_TOKENS = 12_000 +ARGS = ("-t", "file", "--reasoning", "high") + + +def _prompt_tokens(body: dict[str, Any]) -> int: + return 20_000 + len(json.dumps(body.get("messages", []))) // 4 + + +def _summary(_rec: dict[str, Any]) -> Say: + return Say(f"## Goal\nKeep reading the seeded files ({SUMMARY_MARK}).\n## Progress\n- Read several files.\n", chunk_chars=64) + + +def _fake(tmp: Path, name: str, script: list[Any], **kw: Any) -> FakeVertex: + fake = FakeVertex(tmp / name / "fake", project=PROJECT, region=REGION, sa_project=SA_EMBEDDED_PROJECT, + script=script, **kw) + fake.start() + return fake + + +def run_compaction(tmp: Path) -> dict[str, Any]: + script: list[Any] = [Call([("read_file", {"path": f"f{i}.txt"})]) for i in range(4)] + [Say(TURN1_FINAL)] + script += [Call([("read_file", {"path": f"f{i}.txt"})]) for i in range(4, 7)] + [Say(TURN2_FINAL)] + fake = _fake(tmp, "compaction", script, aux=_summary, prompt_tokens_fn=_prompt_tokens) + setup = hermes_setup(fake, context_length=CONTEXT_LENGTH, extra_config={ + "compression": {"threshold_tokens": THRESHOLD_TOKENS, "protect_last_n": 4}}) + nh = make_home(tmp / "compaction" / "h", **setup) + for i in range(7): + (nh.project / f"f{i}.txt").write_text(f"file {i} " + "lorem ipsum dolor " * 300 + "\n", encoding="utf-8") + turn1 = run_chat(nh, "Read f0.txt through f3.txt.", env=fake.child_env(), args=ARGS) + sid = latest_session(nh) if turn1.returncode == 0 else None + turn2 = run_chat(nh, "Now read f4.txt through f6.txt.", env=fake.child_env(), args=ARGS, resume=sid) if sid else None + return {"fake": fake, "nh": nh, "turn1": turn1, "turn2": turn2} + + +def run_drop(tmp: Path) -> dict[str, Any]: + fake = _fake(tmp, "drop", [Call([("read_file", {"path": "note.txt"})]), Drop(DROP_PARTIAL, after_chars=22), Say(DROP_FINAL)]) + nh = make_home(tmp / "drop" / "h", **hermes_setup(fake)) + (nh.project / "note.txt").write_text("note: KIWI-13\n", encoding="utf-8") + return {"fake": fake, "nh": nh, "turn": run_chat(nh, "Read note.txt.", env=fake.child_env(), args=ARGS)} + + +SCENARIOS: dict[str, Callable[[Path], dict[str, Any]]] = {"compaction": run_compaction, "drop": run_drop} + + +@pytest.fixture(scope="module") +def results(tmp_path_factory: pytest.TempPathFactory) -> Any: + tmp = tmp_path_factory.mktemp("vertex_recovery") + with ThreadPoolExecutor(max_workers=len(SCENARIOS)) as pool: + futures = {name: pool.submit(fn, tmp) for name, fn in SCENARIOS.items()} + out = {name: f.result() for name, f in futures.items()} + yield out + for res in out.values(): + res["fake"].stop() + + +def _ok(turn: ChatResult | None, what: str, answer: str) -> None: + assert turn is not None and turn.returncode == 0, f"{what} failed:\n{turn.describe() if turn else 'not run'}" + assert answer in turn.stdout, turn.describe() + + +def _assert_pairs(rows: list[dict[str, Any]]) -> None: + calls = [tc["id"] for r in rows if r["role"] == "assistant" for tc in tool_calls_of(r)] + results = [r["tool_call_id"] for r in rows if r["role"] == "tool"] + assert sorted(calls) == sorted(results), f"persisted tool pairs broken: calls={calls} results={results}" + + +def test_compacted_signed_session_stays_valid_for_vertex(results: dict[str, Any]) -> None: + """The summary call goes to Vertex with the minted bearer, and every request after compaction is + accepted: tool pairs intact, current-turn calls still signed, each signature on its own call.""" + res = results["compaction"] + fake = res["fake"] + _ok(res["turn1"], "turn 1", TURN1_FINAL) + _ok(res["turn2"], "turn 2 (--resume)", TURN2_FINAL) + assert not fake.rejected(), json.dumps([(r["status"], r["rejected"]) for r in fake.rejected()], indent=1) + aux = fake.aux_requests() + assert aux, "compaction never called the summarizer (threshold not crossed?)" + minted = {f"Bearer {t}" for t in fake.minted_tokens()} + assert {r["auth"] for r in aux} <= minted + first_aux = fake.requests.index(aux[0]) + after = [r for r in fake.requests[first_aux:] if r["kind"] == "main"] + assert after and any(SUMMARY_MARK in json.dumps(r["body"]["messages"]) for r in after), ( + "no post-compaction request carries the summary") + for rec in after: + for msg in rec["body"]["messages"]: + for tc in msg.get("tool_calls") or []: + sig = (tc.get("extra_content") or {}).get("google", {}).get("thought_signature") + if sig is not None: + assert fake.signature_by_call.get(tc["id"]) == sig, f"signature moved to another call: {tc['id']}" + + +def test_compacted_session_rows_keep_tool_pairs(results: dict[str, Any]) -> None: + res = results["compaction"] + nh = res["nh"] + _ok(res["turn2"], "turn 2 (--resume)", TURN2_FINAL) + rows = messages(nh, latest_session(nh)) + _assert_pairs(rows) + assert_no_duplicate_assistant_text(rows, TURN2_FINAL) + assert any(r["role"] == "assistant" and r.get("content") == TURN2_FINAL for r in rows) + + +def test_stream_drop_retried_without_duplicate_content(results: dict[str, Any]) -> None: + """The response after a signed tool step dies mid-body: Hermes resends the same valid request + (signature intact), prints the recovered answer, and persists it exactly once.""" + res = results["drop"] + fake, nh = res["fake"], res["nh"] + _ok(res["turn"], "drop turn", DROP_FINAL) + assert not fake.rejected(), [r["rejected"] for r in fake.rejected()] + mains = fake.main_requests() + assert [r["response"] for r in mains] == ["Call", "Drop", "Say"], [r.get("response") for r in mains] + assert mains[2]["body"]["messages"] == mains[1]["body"]["messages"], "the retry changed the conversation" + rows = messages(nh, latest_session(nh)) + assert [r["role"] for r in rows] == ["user", "assistant", "tool", "assistant"], rows + partial = DROP_PARTIAL[:22] + assert not [r["id"] for r in rows if partial in (r.get("content") or "")], "the dropped partial was persisted" + assert_no_duplicate_assistant_text(rows, DROP_FINAL) + assert partial not in res["turn"].stdout diff --git a/tests/e2e/core/providers/test_native_vertex_tools.py b/tests/e2e/core/providers/test_native_vertex_tools.py new file mode 100644 index 000000000000..d63048d83d83 --- /dev/null +++ b/tests/e2e/core/providers/test_native_vertex_tools.py @@ -0,0 +1,247 @@ +"""Vertex AI wire conformance: OAuth minting, multi-turn tools, thought-signature replay after resume. + +The real ``hermes chat -q`` CLI talks to Vertex through the standard ``HTTPS_PROXY`` + +``SSL_CERT_FILE`` channel; the fake (``tests/fakes/providers/vertex.py``) terminates TLS for +``us-central1-aiplatform.googleapis.com``, validates every request against the Vertex +OpenAI-compatibility contract, and mints tokens for the REAL ``google-auth`` JWT exchange. + +Scenarios run concurrently (one fake + one hermetic home each) in a module fixture: + +* ``session`` — turn 1 calls ``read_file`` (signed call), turn 2 runs in a NEW process via + ``--resume`` and calls it again; ``--reasoning high`` throughout. +* ``refresh`` — tokens are minted with a short ``expires_in``; the vendor expires them right + after the first response, so the next request 401s and must be retried with a re-minted token. +* ``default_toolset`` — a turn with Hermes' default toolsets (every tool schema goes to Vertex). +""" + +from __future__ import annotations + +import json +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from typing import Any, Callable + +import pytest + +pytest.importorskip("google.auth", reason="Vertex minting needs google-auth (CI installs it)") + +from tests.e2e.core.providers._native_helpers import ( # noqa: E402 + ChatResult, + KnownSymptom, + NativeHome, + latest_session, + make_home, + messages, + run_chat, + tool_calls_of, +) +from tests.fakes.providers.vertex import ( # noqa: E402 + PROJECT, + REGION, + SA_EMBEDDED_PROJECT, + Call, + FakeVertex, + Say, + hermes_setup, + signatures_on_wire, +) + +# key -> "#issue reason"; a strict xfail turns red the moment the bug is fixed. +KNOWN: dict[str, str] = { + "default_toolset": "#109115 terminal.notify anyOf[boolean, array] is rejected by Vertex's FunctionDeclaration " + "translator, so every default-toolset turn 400s", +} + +SECRET_1 = "PINEAPPLE-42" +SECRET_2 = "MANGO-77" +FINAL_1 = "Turn one answer: the note says PINEAPPLE-42." +FINAL_2 = "Turn two answer: the second note says MANGO-77." +FINAL_REFRESH = "Answer after the token was re-minted." +FILE_ONLY = ("-t", "file") + + +class Precondition(RuntimeError): + """A scenario broke before reaching the property under test (never masquerades as a KNOWN xfail).""" + + +def require(ok: Any, what: str) -> None: + if not ok: + raise Precondition(what) + + +def _home(tmp: Path, name: str, fake: FakeVertex) -> NativeHome: + nh = make_home(tmp / name / "h", **hermes_setup(fake)) + (nh.project / "note1.txt").write_text(f"note: {SECRET_1}\n", encoding="utf-8") + (nh.project / "note2.txt").write_text(f"note: {SECRET_2}\n", encoding="utf-8") + return nh + + +def _fake(tmp: Path, name: str, script: list[Any]) -> FakeVertex: + fake = FakeVertex(tmp / name / "fake", project=PROJECT, region=REGION, sa_project=SA_EMBEDDED_PROJECT, script=script) + fake.start() + return fake + + +def run_session(tmp: Path) -> dict[str, Any]: + fake = _fake(tmp, "session", [ + Call([("read_file", {"path": "note1.txt"})]), Say(FINAL_1), + Call([("read_file", {"path": "note2.txt"})]), Say(FINAL_2), + ]) + nh = _home(tmp, "session", fake) + args = (*FILE_ONLY, "--reasoning", "high") + turn1 = run_chat(nh, "Read note1.txt and tell me what it says.", env=fake.child_env(), args=args) + boundary = len(fake.requests) + tokens_after_turn1 = len(fake.token_requests) + sid = latest_session(nh) if turn1.returncode == 0 else None + turn2 = run_chat(nh, "Now read note2.txt.", env=fake.child_env(), args=args, resume=sid) if sid else None + return {"fake": fake, "nh": nh, "turn1": turn1, "turn2": turn2, "boundary": boundary, + "tokens_after_turn1": tokens_after_turn1, "sid": sid} + + +def run_refresh(tmp: Path) -> dict[str, Any]: + fake = _fake(tmp, "refresh", [Call([("read_file", {"path": "note1.txt"})], expire_tokens_after=True), Say(FINAL_REFRESH)]) + # Below google-auth's refresh window, so a re-mint yields a NEW token rather than the cached one. + fake.token_policy.expires_in = 120 + nh = _home(tmp, "refresh", fake) + turn = run_chat(nh, "Read note1.txt.", env=fake.child_env(), args=FILE_ONLY) + return {"fake": fake, "nh": nh, "turn": turn} + + +def run_default_toolset(tmp: Path) -> dict[str, Any]: + fake = _fake(tmp, "default_toolset", [Say("Default toolset answer.")]) + nh = _home(tmp, "default_toolset", fake) + return {"fake": fake, "nh": nh, "turn": run_chat(nh, "Say hello.", env=fake.child_env())} + + +SCENARIOS: dict[str, Callable[[Path], dict[str, Any]]] = { + "session": run_session, "refresh": run_refresh, "default_toolset": run_default_toolset, +} + + +@pytest.fixture(scope="module") +def results(tmp_path_factory: pytest.TempPathFactory) -> Any: + tmp = tmp_path_factory.mktemp("vertex_tools") + with ThreadPoolExecutor(max_workers=len(SCENARIOS)) as pool: + futures = {name: pool.submit(fn, tmp) for name, fn in SCENARIOS.items()} + out = {name: f.result() for name, f in futures.items()} + yield out + for res in out.values(): + res["fake"].stop() + + +def _ok(turn: ChatResult | None, what: str) -> ChatResult: + require(turn is not None and turn.returncode == 0, f"{what} failed:\n{turn.describe() if turn else 'not run'}") + assert turn is not None + return turn + + +def _rejections(fake: FakeVertex) -> str: + return json.dumps([(r["status"], r["rejected"]) for r in fake.rejected()], indent=1) + + +def test_tool_result_goes_back_paired_to_the_signed_call(results: dict[str, Any]) -> None: + """Turn 1: the model's ``read_file`` call runs for real and its result goes back as a ``tool`` + message paired to the call id, with the call's thought signature replayed byte-for-byte.""" + res = results["session"] + fake, nh = res["fake"], res["nh"] + turn1 = _ok(res["turn1"], "turn 1") + assert FINAL_1 in turn1.stdout, turn1.describe() + assert not fake.rejected(), f"Vertex rejected requests:\n{_rejections(fake)}" + first, second = fake.main_requests()[:2] + (issued,) = first["tool_calls"] + msgs = second["body"]["messages"] + asst = next(m for m in msgs if m.get("role") == "assistant" and m.get("tool_calls")) + (sent,) = asst["tool_calls"] + assert sent["id"] == issued["id"] and sent["function"]["name"] == "read_file" + assert sent["extra_content"] == issued["extra_content"], "thought signature not replayed verbatim" + tool_msgs = [m for m in msgs if m.get("role") == "tool"] + assert [m["tool_call_id"] for m in tool_msgs] == [issued["id"]] + assert SECRET_1 in tool_msgs[0]["content"], "the real read_file result did not reach the model" + rows = messages(nh, res["sid"])[:4] # turn 2 (--resume) appends to the same session afterwards + assert [r["role"] for r in rows] == ["user", "assistant", "tool", "assistant"], rows + persisted = tool_calls_of(rows[1]) + assert [tc["id"] for tc in persisted] == [issued["id"]] + assert rows[2]["tool_call_id"] == issued["id"] and rows[3]["content"] == FINAL_1 + + +def test_wire_scheme_bearer_and_single_token_mint(results: dict[str, Any]) -> None: + """Every request hits the configured project/location path on the regional host with a bearer the + OAuth endpoint minted from a verified SA JWT; one mint per process, reused across API calls.""" + res = results["session"] + fake = res["fake"] + _ok(res["turn1"], "turn 1") + _ok(res["turn2"], "turn 2") + assert all(r["claims"] for r in fake.token_requests), fake.token_requests + assert [c["target"] for c in fake.connects if c["allowed"]], "no request reached Vertex through the proxy" + per_process = [fake.requests[: res["boundary"]], fake.requests[res["boundary"]:]] + minted = fake.minted_tokens() + assert res["tokens_after_turn1"] == 1 and len(minted) == 2, ( + f"expected one token exchange per process, saw {len(fake.token_requests)}: {fake.token_requests}") + for token, reqs in zip(minted, per_process): + assert len(reqs) >= 2 + assert {r["auth"] for r in reqs} == {f"Bearer {token}"}, "bearer not the minted token / not reused" + expected_path = f"/v1beta1/projects/{PROJECT}/locations/{REGION}/endpoints/openapi/chat/completions" + assert {(r["host"], r["path"]) for r in fake.requests} == {(f"{REGION}-aiplatform.googleapis.com", expected_path)} + assert {r["body"]["model"] for r in fake.requests} == {"google/gemini-3-flash-preview"} + + +def test_reasoning_effort_reaches_vertex_as_thinking_config(results: dict[str, Any]) -> None: + """``--reasoning high`` rides in the documented ``extra_body.google.thinking_config`` wrapper + (a bare top-level ``google`` key is silently ignored by Vertex) and never alongside + ``reasoning_effort`` (Vertex allows only one of the two).""" + res = results["session"] + _ok(res["turn1"], "turn 1") + for rec in res["fake"].main_requests(): + body = rec["body"] + thinking = body.get("extra_body", {}).get("google", {}).get("thinking_config") + assert thinking and thinking.get("thinking_level") == "high", {k: v for k, v in body.items() if k != "messages"} + assert "reasoning_effort" not in body + + +def test_thought_signature_replayed_after_resume_in_new_process(results: dict[str, Any]) -> None: + """Turn 2 runs in a fresh process: its first request replays turn 1's signed call (same id, same + signature bytes) from state.db, and Vertex accepts every request of the resumed turn.""" + res = results["session"] + fake = res["fake"] + turn2 = _ok(res["turn2"], "turn 2 (--resume)") + assert FINAL_2 in turn2.stdout, turn2.describe() + assert not fake.rejected(), f"Vertex rejected requests:\n{_rejections(fake)}" + turn1_call = fake.main_requests()[0]["tool_calls"][0] + resumed = fake.requests[res["boundary"]] + replayed = {tc["id"]: tc for m in resumed["body"]["messages"] if m.get("role") == "assistant" + for tc in m.get("tool_calls") or []} + assert turn1_call["id"] in replayed, "turn 1's tool call missing from the resumed request" + assert replayed[turn1_call["id"]].get("extra_content") == turn1_call["extra_content"] + final = fake.main_requests()[-1]["body"] + assert signatures_on_wire(final) == fake.issued_signatures, "resumed turn lost or reordered signatures" + assert SECRET_2 in json.dumps(final["messages"]) + + +def test_expired_token_is_reminted_and_request_retried(results: dict[str, Any]) -> None: + """The vendor expires the bearer mid-turn: the 401 UNAUTHENTICATED triggers a fresh JWT exchange + and the SAME request is retried once with the new token; the user only sees the answer.""" + res = results["refresh"] + fake, nh = res["fake"], res["nh"] + turn = _ok(res["turn"], "refresh turn") + assert FINAL_REFRESH in turn.stdout, turn.describe() + statuses = [r.get("status") for r in fake.requests] + assert statuses == [200, 401, 200], statuses + stale, retried = fake.requests[1], fake.requests[2] + assert stale["auth"] == fake.requests[0]["auth"] and retried["auth"] != stale["auth"] + assert retried["auth"] == f"Bearer {fake.minted_tokens()[-1]}", "retry did not use the newest minted token" + assert retried["body"]["messages"] == stale["body"]["messages"] + assert "401" not in turn.stdout and "UNAUTHENTICATED" not in turn.stdout + rows = messages(nh, latest_session(nh)) + assert [r["role"] for r in rows] == ["user", "assistant", "tool", "assistant"], rows + + +@pytest.mark.xfail(strict=True, raises=KnownSymptom, reason=KNOWN["default_toolset"]) +def test_default_toolset_schemas_accepted_by_vertex(results: dict[str, Any]) -> None: + """With Hermes' default toolsets every tool declaration must survive Vertex's translation.""" + res = results["default_toolset"] + fake = res["fake"] + require(fake.requests and fake.requests[0]["auth"].startswith("Bearer ya29."), "turn never reached Vertex") + schema_rejects = [r["rejected"] for r in fake.rejected() if "schema type should be ARRAY" in (r["rejected"] or "")] + if schema_rejects: + raise KnownSymptom(f"Vertex rejected the tool declarations: {schema_rejects[0]}") + assert "Default toolset answer." in res["turn"].stdout, res["turn"].describe() diff --git a/tests/fakes/providers/__init__.py b/tests/fakes/providers/__init__.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/tests/fakes/providers/bedrock_converse.py b/tests/fakes/providers/bedrock_converse.py new file mode 100644 index 000000000000..9f98ec5f4bf7 --- /dev/null +++ b/tests/fakes/providers/bedrock_converse.py @@ -0,0 +1,499 @@ +"""Loopback fake of the AWS Bedrock Runtime Converse / ConverseStream API. + +The real ``boto3`` ``bedrock-runtime`` client inside Hermes is pointed here through botocore's +documented endpoint override (``AWS_ENDPOINT_URL_BEDROCK_RUNTIME``), so the SDK serializes, signs +(SigV4) and parses exactly as against AWS; only the service is fake. + +What the fake enforces, the way Bedrock does: + +* SigV4: every request must carry ``Authorization: AWS4-HMAC-SHA256 Credential=/// + bedrock/aws4_request, SignedHeaders=..., Signature=...``; the signature is RE-COMPUTED with the known + fake secret and must match (a tampered body or wrong key is a 403, like AWS). +* Request shape: the JSON body plus the ``modelId`` from the URI is validated against the botocore + service model's ``Converse`` / ``ConverseStream`` input shape (``ParamValidator``: types, required + members, tagged unions). Violations are a 400 ``ValidationException``. +* Conversation semantics the service model cannot express but Converse rejects (messages the real + endpoint answers with a ValidationException): first message must be ``user``; roles alternate; + text blocks are non-blank; every ``toolResult`` answers a ``toolUse`` of the immediately + preceding assistant turn and every such ``toolUse`` is answered; a replayed ``reasoningText`` must + carry exactly the signature this fake issued for that text (unsigned or altered thinking is + rejected, as signed-thinking models do). + +Errors use the AWS JSON error shape (``x-amzn-ErrorType`` header + ``{"message": ...}``); +ConverseStream answers real ``application/vnd.amazon.eventstream`` binary frames (prelude, typed +headers, CRC32s), including ``:message-type exception`` frames and mid-stream connection drops. +""" + +from __future__ import annotations + +import base64 +import binascii +import json +import re +import struct +import threading +import time +from dataclasses import dataclass, field +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import Any, Callable, Union +from urllib.parse import unquote + +import botocore.session +from botocore.auth import SigV4Auth +from botocore.awsrequest import AWSRequest +from botocore.credentials import Credentials +from botocore.validate import ParamValidator + +ACCESS_KEY = "AKIAFAKEE2EBEDROCK01" +SECRET_KEY = "fake/e2e/secret/never+real/0000000000000" +REGION = "us-east-1" + +_SERVICE = botocore.session.get_session().get_service_model("bedrock-runtime") +_ROUTE_RE = re.compile(r"^/model/(?P[^/]+)/(?Pconverse|converse-stream)$") +_OPS = {"converse": "Converse", "converse-stream": "ConverseStream"} +_AUTH_RE = re.compile( + r"^AWS4-HMAC-SHA256 Credential=(?P[^/]+)/(?P\d{8})/(?P[a-z0-9-]+)/" + r"(?P[a-z0-9-]+)/aws4_request, ?SignedHeaders=(?P[a-z0-9;-]+), ?" + r"Signature=(?P[0-9a-f]{64})$") +# HTTP status + x-amzn-ErrorType per Bedrock Runtime error (botocore service model metadata). +ERROR_STATUS = { + name: shape.metadata["error"]["httpStatusCode"] + for name in _SERVICE.shape_names + if (shape := _SERVICE.shape_for(name)).metadata.get("exception") +} + + +# -------------------------------------------------------------------------------------------------- +# Scripted content (Converse ContentBlock shapes) +# -------------------------------------------------------------------------------------------------- + + +@dataclass +class Text: + text: str + chunk: int = 12 + + +@dataclass +class Reasoning: + """``reasoningContent.reasoningText``; the signature is issued by the fake and remembered.""" + + text: str + signature: str = "" + + +@dataclass +class ToolUse: + name: str + input: dict[str, Any] + tool_use_id: str = "" + + +Block = Union[Text, Reasoning, ToolUse] + + +@dataclass +class Turn: + """One assistant message (streamed or not) + usage the service reports.""" + + blocks: list[Block] + stop_reason: str = "" + input_tokens: int = 120 + output_tokens: int = 30 + + +@dataclass +class HttpError: + """An AWS JSON error response before any stream starts (e.g. 429 ThrottlingException).""" + + code: str + message: str + + +@dataclass +class StreamException: + """ConverseStream: emit ``after`` events of ``turn`` then a ``:message-type exception`` frame.""" + + turn: Turn + exception: str # event member name, e.g. "throttlingException" + message: str + after: int = 2 + + +@dataclass +class Drop: + """ConverseStream: emit ``after`` events of ``turn`` then close the socket mid-stream. + + ``clean=True`` ends the HTTP body properly instead (valid framing, but the event stream stops + before ``messageStop``: the service never finished the message).""" + + turn: Turn + after: int = 3 + clean: bool = False + + +Reply = Union[Turn, HttpError, StreamException, Drop] + + +# -------------------------------------------------------------------------------------------------- +# Event-stream framing (application/vnd.amazon.eventstream) +# -------------------------------------------------------------------------------------------------- + + +def _header(name: str, value: str) -> bytes: + n, v = name.encode(), value.encode() + return struct.pack("!B", len(n)) + n + b"\x07" + struct.pack("!H", len(v)) + v # 7 = string + + +def encode_frame(headers: dict[str, str], payload: bytes) -> bytes: + """One event-stream message: prelude(total, headers_len) + prelude CRC + headers + payload + CRC.""" + hdr = b"".join(_header(k, v) for k, v in headers.items()) + total = 12 + len(hdr) + len(payload) + 4 + prelude = struct.pack("!II", total, len(hdr)) + head = prelude + struct.pack("!I", binascii.crc32(prelude) & 0xFFFFFFFF) + hdr + payload + return head + struct.pack("!I", binascii.crc32(head) & 0xFFFFFFFF) + + +def event_frame(event_type: str, body: dict[str, Any]) -> bytes: + _validate_event(event_type, body) + return encode_frame({":event-type": event_type, ":content-type": "application/json", + ":message-type": "event"}, json.dumps(body).encode()) + + +def exception_frame(exception: str, message: str) -> bytes: + return encode_frame({":exception-type": exception, ":content-type": "application/json", + ":message-type": "exception"}, json.dumps({"message": message}).encode()) + + +def _validate_event(event_type: str, body: dict[str, Any]) -> None: + """The fake only emits events valid for the ConverseStream output shape (self-check).""" + shape = _SERVICE.operation_model("ConverseStream").output_shape.members["stream"].members[event_type] + report = ParamValidator().validate(body, shape) + if report.has_errors(): + raise AssertionError(f"fake built an invalid {event_type} event: {report.generate_report()}") + + +# -------------------------------------------------------------------------------------------------- +# Request validation +# -------------------------------------------------------------------------------------------------- + + +class Rejection(Exception): + def __init__(self, code: str, message: str) -> None: + super().__init__(message) + self.code = code + self.message = message + + +def _wire_to_params(shape: Any, value: Any) -> Any: + """rest-json wire value -> botocore param value (blobs arrive base64-encoded).""" + kind = shape.type_name + if kind == "blob" and isinstance(value, str): + try: + return base64.b64decode(value, validate=True) + except (binascii.Error, ValueError): + raise Rejection("ValidationException", f"Invalid base64 blob for {shape.name}") from None + if kind == "structure" and isinstance(value, dict) and not shape.metadata.get("document"): + return {k: (_wire_to_params(shape.members[k], v) if k in shape.members else v) for k, v in value.items()} + if kind == "list" and isinstance(value, list): + return [_wire_to_params(shape.member, v) for v in value] + if kind == "map" and isinstance(value, dict): + return {k: _wire_to_params(shape.value, v) for k, v in value.items()} + return value + + +def validate_shape(op: str, model_id: str, body: dict[str, Any]) -> None: + shape = _SERVICE.operation_model(op).input_shape + params = _wire_to_params(shape, {**body, "modelId": model_id}) + report = ParamValidator().validate(params, shape) + if report.has_errors(): + raise Rejection("ValidationException", report.generate_report()) + + +def _tool_ids(content: list[dict[str, Any]], kind: str) -> list[str]: + return [b[kind]["toolUseId"] for b in content if kind in b] + + +def validate_conversation(body: dict[str, Any], signatures: dict[str, str], reasoned_tools: set[str]) -> None: + """Converse-side rules beyond the schema (real ValidationException texts).""" + messages = body.get("messages") or [] + if not messages or messages[0]["role"] != "user": + raise Rejection("ValidationException", "A conversation must start with a user message. " + "Try again with a conversation that starts with a user message.") + for i, msg in enumerate(messages): + if i and msg["role"] == messages[i - 1]["role"]: + raise Rejection("ValidationException", "A conversation must alternate between user and " + "assistant roles. Make sure the conversation alternates between user and " + "assistant roles and try again.") + for j, block in enumerate(msg["content"]): + if "text" in block and not block["text"].strip(): + raise Rejection("ValidationException", f"The text field in the ContentBlock object at " + f"messages.{i}.content.{j} is blank. Add text to the text field, and try again.") + _check_reasoning(block.get("reasoningContent"), signatures, f"messages.{i}.content.{j}") + expected = _tool_ids(messages[i - 1]["content"], "toolUse") if i and msg["role"] == "user" else [] + got = _tool_ids(msg["content"], "toolResult") + if sorted(expected) != sorted(got): + raise Rejection("ValidationException", f"Expected toolResult blocks at messages.{i}.content " + f"for the following Ids: {', '.join(expected) or '(none)'}; got {', '.join(got) or '(none)'}") + _check_final_assistant_thinking(messages, reasoned_tools) + + +def _check_final_assistant_thinking(messages: list[dict[str, Any]], reasoned_tools: set[str]) -> None: + """Signed-thinking models: while a tool loop is open (the request ends in toolResults), the final + assistant turn must START with its reasoning block if it was issued with one.""" + if len(messages) < 2 or not _tool_ids(messages[-1]["content"], "toolResult"): + return + final = messages[-2]["content"] + if set(_tool_ids(final, "toolUse")) & reasoned_tools and "reasoningContent" not in final[0]: + first = next(iter(final[0]), "?") + raise Rejection("ValidationException", f"messages.{len(messages) - 2}.content.0.type: Expected " + f"`thinking` or `redacted_thinking`, but found `{first}`. When `thinking` is enabled, " + "a final `assistant` message must start with a thinking block (preceeding the lastmost " + "set of `tool_use` and `tool_result` blocks).") + + +def _check_reasoning(reasoning: Any, signatures: dict[str, str], where: str) -> None: + text_block = (reasoning or {}).get("reasoningText") + if not text_block: + return + issued = signatures.get(text_block.get("text", "")) + if issued is None or text_block.get("signature") != issued: + raise Rejection("ValidationException", f"{where}: The reasoning block signature is missing or " + "invalid. Reasoning content must be passed back unmodified with its signature.") + + +def verify_sigv4(method: str, url: str, headers: dict[str, str], body: bytes) -> dict[str, str]: + """Re-compute the SigV4 signature with the fake secret; return the parsed credential scope.""" + auth = headers.get("authorization", "") + m = _AUTH_RE.match(auth) + if not m: + raise Rejection("MissingAuthenticationTokenException", f"Missing or malformed SigV4 Authorization: {auth!r}") + if m["key"] != ACCESS_KEY or m["service"] != "bedrock": + raise Rejection("UnrecognizedClientException", "The security token included in the request is invalid.") + signed = {h: headers[h] for h in m["signed"].split(";") if h in headers} + request = AWSRequest(method=method, url=url, data=body, headers=signed) + request.context["timestamp"] = headers.get("x-amz-date", "") + signer = SigV4Auth(Credentials(ACCESS_KEY, SECRET_KEY), "bedrock", m["region"]) + canonical = signer.canonical_request(request) + expected = signer.signature(signer.string_to_sign(request, canonical), request) + if expected != m["sig"]: + raise Rejection("InvalidSignatureException", "The request signature we calculated does not match " + "the signature you provided.") + return m.groupdict() + + +# -------------------------------------------------------------------------------------------------- +# Server +# -------------------------------------------------------------------------------------------------- + + +Responder = Callable[[dict[str, Any]], Reply] + + +@dataclass +class FakeBedrock: + """Recording, validating Bedrock Runtime fake. ``responder(record) -> Reply`` scripts each call.""" + + responder: Responder + requests: list[dict[str, Any]] = field(default_factory=list) + signatures: dict[str, str] = field(default_factory=dict) # reasoning text -> issued signature + reasoned_tools: set[str] = field(default_factory=set) # toolUseIds issued in a turn with reasoning + _lock: threading.Lock = field(default_factory=threading.Lock) + _ids: int = 0 + _httpd: ThreadingHTTPServer | None = None + + def __enter__(self) -> "FakeBedrock": + self._httpd = ThreadingHTTPServer(("127.0.0.1", 0), _handler_for(self)) + self._httpd.daemon_threads = True + threading.Thread(target=self._httpd.serve_forever, daemon=True).start() + return self + + def __exit__(self, *_exc: object) -> None: + if self._httpd: + self._httpd.shutdown() + self._httpd.server_close() + + @property + def endpoint(self) -> str: + assert self._httpd is not None + return f"http://127.0.0.1:{self._httpd.server_address[1]}" + + def client_env(self) -> dict[str, str]: + """Env for a Hermes child: botocore endpoint override + fake static credentials.""" + return {"AWS_ENDPOINT_URL_BEDROCK_RUNTIME": self.endpoint, "AWS_ACCESS_KEY_ID": ACCESS_KEY, + "AWS_SECRET_ACCESS_KEY": SECRET_KEY, "AWS_REGION": REGION} + + def snapshot(self) -> list[dict[str, Any]]: + with self._lock: + return list(self.requests) + + def next_id(self, prefix: str) -> str: + with self._lock: + self._ids += 1 + return f"{prefix}{self._ids:04d}{int(time.monotonic() * 1000) % 100000:05d}" + + def record(self, rec: dict[str, Any]) -> None: + with self._lock: + self.requests.append(rec) + + def materialize(self, turn: Turn) -> list[dict[str, Any]]: + """Resolve ids/signatures and return Converse ContentBlocks for the turn.""" + out: list[dict[str, Any]] = [] + for block in turn.blocks: + if isinstance(block, Text): + out.append({"text": block.text}) + elif isinstance(block, Reasoning): + block.signature = block.signature or base64.b64encode( + f"sig:{self.next_id('r')}:{len(block.text)}".encode()).decode() + with self._lock: + self.signatures[block.text] = block.signature + out.append({"reasoningContent": {"reasoningText": {"text": block.text, "signature": block.signature}}}) + else: + block.tool_use_id = block.tool_use_id or self.next_id("tooluse_") + out.append({"toolUse": {"toolUseId": block.tool_use_id, "name": block.name, "input": block.input}}) + if any(isinstance(b, Reasoning) for b in turn.blocks): + with self._lock: + self.reasoned_tools.update(b.tool_use_id for b in turn.blocks if isinstance(b, ToolUse)) + if not turn.stop_reason: + turn.stop_reason = "tool_use" if any(isinstance(b, ToolUse) for b in turn.blocks) else "end_turn" + return out + + +def _usage(turn: Turn) -> dict[str, int]: + return {"inputTokens": turn.input_tokens, "outputTokens": turn.output_tokens, + "totalTokens": turn.input_tokens + turn.output_tokens} + + +def _pieces(text: str, size: int) -> list[str]: + return [text[i:i + size] for i in range(0, len(text), size)] or [""] + + +def _block_events(index: int, block: dict[str, Any], chunk: int) -> list[tuple[str, dict[str, Any]]]: + """ContentBlockStart/Delta/Stop events for one block (text blocks get no start, as on AWS).""" + events: list[tuple[str, dict[str, Any]]] = [] + if "toolUse" in block: + tu = block["toolUse"] + events.append(("contentBlockStart", {"contentBlockIndex": index, "start": { + "toolUse": {"toolUseId": tu["toolUseId"], "name": tu["name"]}}})) + deltas = [{"toolUse": {"input": p}} for p in _pieces(json.dumps(tu["input"]), chunk)] + elif "reasoningContent" in block: + rt = block["reasoningContent"]["reasoningText"] + deltas = [{"reasoningContent": {"text": p}} for p in _pieces(rt["text"], chunk)] + deltas.append({"reasoningContent": {"signature": rt["signature"]}}) + else: + deltas = [{"text": p} for p in _pieces(block["text"], chunk)] + events += [("contentBlockDelta", {"contentBlockIndex": index, "delta": d}) for d in deltas] + events.append(("contentBlockStop", {"contentBlockIndex": index})) + return events + + +def stream_events(blocks: list[dict[str, Any]], turn: Turn) -> list[bytes]: + chunk = min((b.chunk for b in turn.blocks if isinstance(b, Text)), default=12) + events: list[tuple[str, dict[str, Any]]] = [("messageStart", {"role": "assistant"})] + for i, block in enumerate(blocks): + events += _block_events(i, block, chunk) + events.append(("messageStop", {"stopReason": turn.stop_reason})) + events.append(("metadata", {"usage": _usage(turn), "metrics": {"latencyMs": 7}})) + return [event_frame(name, body) for name, body in events] + + +def _handler_for(fake: FakeBedrock) -> type[BaseHTTPRequestHandler]: + class Handler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def log_message(self, *_args: Any) -> None: + return + + def do_POST(self) -> None: # noqa: N802 - http.server API + raw = self.rfile.read(int(self.headers.get("Content-Length") or 0)) + headers = {k.lower(): v for k, v in self.headers.items()} + route = _ROUTE_RE.match(self.path.split("?", 1)[0]) + rec: dict[str, Any] = {"path": self.path, "headers": headers, "raw": raw, "t": time.time(), + "op": _OPS.get(route["op"]) if route else None, + "model": unquote(route["model"]) if route else None} + try: + body = json.loads(raw or b"{}") + except ValueError: + body = None + rec["body"] = body + try: + if not route: + raise Rejection("UnknownOperationException", f"No route for {self.path}") + rec["auth"] = verify_sigv4("POST", f"http://{headers.get('host')}{self.path}", headers, raw) + if not isinstance(body, dict): + raise Rejection("SerializationException", "Request body is not a JSON object") + validate_shape(rec["op"], rec["model"], body) + validate_conversation(body, dict(fake.signatures), set(fake.reasoned_tools)) + except Rejection as rej: + rec["rejected"] = f"{rej.code}: {rej.message}" + fake.record(rec) + return self._error(rej.code, rej.message) + fake.record(rec) + reply = fake.responder(rec) + rec["reply"] = type(reply).__name__ + if isinstance(reply, HttpError): + return self._error(reply.code, reply.message) + turn = reply if isinstance(reply, Turn) else reply.turn + rec["emitted"] = fake.materialize(turn) + if rec["op"] == "Converse": + return self._converse(turn, rec["emitted"]) + return self._stream(reply, turn, rec["emitted"]) + + def _send(self, status: int, headers: dict[str, str], body: bytes) -> None: + self.send_response(status) + for k, v in {**headers, "Content-Length": str(len(body)), + "x-amzn-RequestId": fake.next_id("req-")}.items(): + self.send_header(k, v) + self.end_headers() + self.wfile.write(body) + self.wfile.flush() + + def _error(self, code: str, message: str) -> None: + status = ERROR_STATUS.get(code, 403 if "Signature" in code or "Token" in code else 400) + self._send(status, {"Content-Type": "application/json", "x-amzn-ErrorType": f"{code}:http://internal.amazon.com/coral/com.amazon.bedrock/"}, + json.dumps({"message": message}).encode()) + + def _converse(self, turn: Turn, blocks: list[dict[str, Any]]) -> None: + body = {"output": {"message": {"role": "assistant", "content": blocks}}, + "stopReason": turn.stop_reason, "usage": _usage(turn), "metrics": {"latencyMs": 9}} + self._send(200, {"Content-Type": "application/json"}, json.dumps(body).encode()) + + def _stream(self, reply: Reply, turn: Turn, blocks: list[dict[str, Any]]) -> None: + frames = stream_events(blocks, turn) + cut = len(frames) if isinstance(reply, Turn) else min(reply.after, len(frames)) + self.send_response(200) + self.send_header("Content-Type", "application/vnd.amazon.eventstream") + self.send_header("Transfer-Encoding", "chunked") + self.send_header("x-amzn-RequestId", fake.next_id("req-")) + self.end_headers() + for frame in frames[:cut]: + self._chunk(frame) + if isinstance(reply, Drop) and not reply.clean: + self.close_connection = True + self.wfile.flush() + self.connection.shutdown(2) # socket.SHUT_RDWR: no terminating chunk + return + if isinstance(reply, StreamException): + self._chunk(exception_frame(reply.exception, reply.message)) + self.wfile.write(b"0\r\n\r\n") + self.wfile.flush() + + def _chunk(self, data: bytes) -> None: + self.wfile.write(f"{len(data):x}\r\n".encode() + data + b"\r\n") + self.wfile.flush() + + return Handler + + +def seq(*replies: Reply | Callable[[dict[str, Any]], Reply]) -> Responder: + """Responder answering the Nth call with ``replies[N]`` (the last one repeats); callables get the record.""" + calls: list[int] = [] + lock = threading.Lock() + + def respond(rec: dict[str, Any]) -> Reply: + with lock: + calls.append(1) + reply = replies[min(len(calls), len(replies)) - 1] + return reply(rec) if callable(reply) else reply + + return respond diff --git a/tests/fakes/providers/codex_app_server.py b/tests/fakes/providers/codex_app_server.py new file mode 100644 index 000000000000..2f756027256e --- /dev/null +++ b/tests/fakes/providers/codex_app_server.py @@ -0,0 +1,755 @@ +"""Fake ``codex app-server``: newline-delimited JSON-RPC over stdio, validated against the protocol. + +Run as a script (``python codex_app_server.py --state-dir DIR app-server``) it plays the codex side of +the wire that ``agent/transports/codex_app_server.py`` speaks. Imported, it offers the test-side +harness (:class:`FakeCodex`) that installs an executable wrapper and reads the recorded transcript. + +Protocol truth is the schema bundle emitted by ``codex app-server generate-json-schema`` (codex-cli +0.147): request params are checked field by field with the real server's serde error strings +(``-32600 "Invalid request: missing field `threadId`"``). The real server IGNORES unknown fields, so +the fake answers them normally but records each one as ``ignored`` — a field codex silently drops is +Hermes intent that never reaches the model, and the suite asserts none are sent. Responses Hermes +gives to server-initiated requests (approvals, elicitation) are validated the same way. + +State lives in ``DIR``: ``scenario.json`` (scripted turns, consumed one per ``turn/start`` across +processes), ``threads.json`` (the "rollout store" that ``thread/resume`` reads, so a NEW Hermes process +can resume a thread) and ``transcript.jsonl`` (every message in both directions, plus spawn/exit +events with PIDs). Only stdlib: the wrapper runs it with the test interpreter. +""" + +from __future__ import annotations + +import contextlib +import json +import os +import signal +import subprocess +import sys +import threading +import time +import uuid +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Callable, Optional + +CLI_VERSION = "0.147.0" +INVALID_REQUEST = -32600 +GRANDCHILD_RELEASE = "release-grandchildren" + + +# --------------------------------------------------------------------------------------------------- +# Protocol schema (subset of the codex app-server v2 bundle Hermes can reach) + serde-style validator +# --------------------------------------------------------------------------------------------------- + +class Invalid(Exception): + """A params/response value the real server's deserializer would reject.""" + + +def _kind(value: Any) -> str: + table = ((bool, "boolean"), (int, "integer"), (float, "floating point"), (str, "string"), + (list, "sequence"), (dict, "map"), (type(None), "null")) + return next((name for typ, name in table if isinstance(value, typ)), "value") + + +def _unexpected(value: Any) -> str: + return f"{_kind(value)} `{value}`" if isinstance(value, (bool, int, float, str)) else _kind(value) + + +@dataclass +class Spec: + check: Callable[[Any, str, list], None] + + def __call__(self, value: Any, path: str, ignored: list) -> None: + self.check(value, path, ignored) + + +def _scalar(py: tuple, expected: str) -> Spec: + def check(v: Any, path: str, ignored: list) -> None: + if not isinstance(v, py) or (bool not in py and isinstance(v, bool)): + raise Invalid(f"invalid type: {_unexpected(v)}, expected {expected}") + return Spec(check) + + +STR, BOOL, INT = _scalar((str,), "a string"), _scalar((bool,), "a boolean"), _scalar((int,), "i64") +ANY = Spec(lambda v, p, i: None) + + +def nullable(spec: Spec) -> Spec: + return Spec(lambda v, p, i: None if v is None else spec(v, p, i)) + + +def enum(*values: str) -> Spec: + def check(v: Any, path: str, ignored: list) -> None: + if not isinstance(v, str): + raise Invalid(f"invalid type: {_unexpected(v)}, expected a string") + if v not in values: + raise Invalid(f"unknown variant `{v}`, expected one of {', '.join(f'`{x}`' for x in values)}") + return Spec(check) + + +def array(item: Spec) -> Spec: + def check(v: Any, path: str, ignored: list) -> None: + if not isinstance(v, list): + raise Invalid(f"invalid type: {_unexpected(v)}, expected a sequence") + for n, element in enumerate(v): + item(element, f"{path}[{n}]", ignored) + return Spec(check) + + +def obj(required: Optional[dict] = None, optional: Optional[dict] = None, *, open_map: bool = False) -> Spec: + required, optional = required or {}, optional or {} + + def check(v: Any, path: str, ignored: list) -> None: + if not isinstance(v, dict): + raise Invalid(f"invalid type: {_unexpected(v)}, expected struct") + for name in required: + if name not in v: + raise Invalid(f"missing field `{name}`") + for name, value in v.items(): + spec = required.get(name) or optional.get(name) + if spec is not None: + spec(value, f"{path}.{name}", ignored) + elif not open_map: + ignored.append(f"{path}.{name}") + return Spec(check) + + +def tagged(tag: str, variants: dict[str, Spec]) -> Spec: + """serde internally tagged enum (``{"type": "text", ...}``).""" + def check(v: Any, path: str, ignored: list) -> None: + if not isinstance(v, dict): + raise Invalid(f"invalid type: {_unexpected(v)}, expected internally tagged enum") + if tag not in v: + raise Invalid(f"missing field `{tag}`") + variant = variants.get(v[tag]) if isinstance(v[tag], str) else None + if variant is None: + raise Invalid(f"unknown variant `{v[tag]}`, expected one of {', '.join(f'`{x}`' for x in variants)}") + variant({k: x for k, x in v.items() if k != tag}, path, ignored) + return Spec(check) + + +def one_of(*specs: Spec) -> Spec: + def check(v: Any, path: str, ignored: list) -> None: + errors = [] + for spec in specs: + try: + spec(v, path, []) + return + except Invalid as exc: + errors.append(str(exc)) + raise Invalid(f"data did not match any variant of untagged enum ({'; '.join(errors)})") + return Spec(check) + + +_TEXT_ELEMENT = obj({"byteRange": obj({"start": INT, "end": INT})}, {"placeholder": nullable(STR)}) +_IMAGE_DETAIL = nullable(enum("auto", "low", "high", "original")) +USER_INPUT = tagged("type", { + "text": obj({"text": STR}, {"text_elements": array(_TEXT_ELEMENT)}), + "image": obj({"url": STR}, {"detail": _IMAGE_DETAIL}), + "localImage": obj({"path": STR}, {"detail": _IMAGE_DETAIL}), + "audio": obj({"url": STR}), "localAudio": obj({"path": STR}), + "skill": obj({"name": STR, "path": STR}), "mention": obj({"name": STR, "path": STR}), +}) +_ASK_FOR_APPROVAL = nullable(one_of(enum("untrusted", "on-request", "never"), obj({"granular": ANY}))) +_THREAD_SETTINGS = { + "approvalPolicy": _ASK_FOR_APPROVAL, + "approvalsReviewer": nullable(enum("user", "auto_review", "guardian_subagent")), + "baseInstructions": nullable(STR), "config": nullable(obj(open_map=True)), "cwd": nullable(STR), + "developerInstructions": nullable(STR), "model": nullable(STR), "modelProvider": nullable(STR), + "personality": nullable(enum("none", "friendly", "pragmatic")), + "sandbox": nullable(enum("read-only", "workspace-write", "danger-full-access")), "serviceTier": nullable(STR), +} +_SANDBOX_POLICY = tagged("type", { + "dangerFullAccess": obj(), "readOnly": obj(optional={"networkAccess": BOOL}), + "externalSandbox": obj(optional={"networkAccess": ANY}), + "workspaceWrite": obj(optional={"excludeSlashTmp": BOOL, "excludeTmpdirEnvVar": BOOL, "networkAccess": BOOL, + "writableRoots": array(STR)}), +}) + +REQUEST_PARAMS: dict[str, Spec] = { + "initialize": obj({"clientInfo": obj({"name": STR, "version": STR}, {"title": nullable(STR)})}, { + "capabilities": nullable(obj(optional={ + "experimentalApi": BOOL, "extensions": nullable(obj(open_map=True)), "mcpServerOpenaiFormElicitation": BOOL, + "optOutNotificationMethods": nullable(array(STR)), "requestAttestation": BOOL})), + }), + "thread/start": obj(optional={**_THREAD_SETTINGS, "ephemeral": nullable(BOOL), "serviceName": nullable(STR), + "sessionStartSource": nullable(enum("startup", "clear")), + "threadSource": nullable(STR)}), + "thread/resume": obj({"threadId": STR}, _THREAD_SETTINGS), + "turn/start": obj({"threadId": STR, "input": array(USER_INPUT)}, { + "approvalPolicy": _ASK_FOR_APPROVAL, "approvalsReviewer": _THREAD_SETTINGS["approvalsReviewer"], + "clientUserMessageId": nullable(STR), "cwd": nullable(STR), "effort": nullable(STR), "model": nullable(STR), + "outputSchema": ANY, "personality": _THREAD_SETTINGS["personality"], "sandboxPolicy": nullable(_SANDBOX_POLICY), + "serviceTier": nullable(STR), "summary": nullable(enum("auto", "concise", "detailed", "none")), + }), + "turn/interrupt": obj({"threadId": STR, "turnId": STR}), + "turn/steer": obj({"threadId": STR, "input": array(USER_INPUT), "expectedTurnId": STR}, + {"clientUserMessageId": nullable(STR)}), + "thread/compact/start": obj({"threadId": STR}), +} + +_EXEC_DECISION = one_of(enum("accept", "acceptForSession", "decline", "cancel"), + obj({"acceptWithExecpolicyAmendment": obj({"execpolicy_amendment": array(STR)})}), + obj({"applyNetworkPolicyAmendment": obj({"network_policy_amendment": obj( + {"action": enum("allow", "deny"), "host": STR})})})) +SERVER_REQUEST_RESULTS: dict[str, Spec] = { + "item/commandExecution/requestApproval": obj({"decision": _EXEC_DECISION}), + "item/fileChange/requestApproval": obj({"decision": enum("accept", "acceptForSession", "decline", "cancel")}), + "item/permissions/requestApproval": obj({"permissions": obj(optional={ + "fileSystem": nullable(obj(open_map=True)), "network": nullable(obj({}, {"enabled": nullable(BOOL)}))})}, + {"scope": enum("turn", "session"), "strictAutoReview": nullable(BOOL)}), + "mcpServer/elicitation/request": obj({"action": enum("accept", "decline", "cancel")}, + {"content": ANY, "_meta": ANY}), +} + + +def validate(spec: Spec, value: Any) -> tuple[Optional[str], list[str]]: + """``(serde error or None, ignored field paths)``.""" + ignored: list[str] = [] + try: + spec(value, "$", ignored) + except Invalid as exc: + return str(exc), ignored + return None, ignored + + +# --------------------------------------------------------------------------------------------------- +# The server +# --------------------------------------------------------------------------------------------------- + +def _now_ms() -> int: + return int(time.time() * 1000) + + +class _Store: + """``threads.json``: thread rollouts + the global scripted-turn cursor, shared across processes.""" + + def __init__(self, path: Path) -> None: + self.path = path + + def load(self) -> dict: + if not self.path.exists(): + return {"threads": {}, "turn_cursor": 0} + return json.loads(self.path.read_text(encoding="utf-8")) + + def save(self, data: dict) -> None: + tmp = self.path.with_suffix(".tmp") + tmp.write_text(json.dumps(data, indent=1), encoding="utf-8") + os.replace(tmp, self.path) + + +class _TurnEnded(Exception): + """A step ended the turn (failure, crash handled, interrupt).""" + + +class FakeAppServer: + def __init__(self, state_dir: Path, argv: list[str]) -> None: + self.state_dir = state_dir + self.scenario = json.loads((state_dir / "scenario.json").read_text(encoding="utf-8")) + self.store = _Store(state_dir / "threads.json") + self.transcript_path = state_dir / "transcript.jsonl" + self._write_lock = threading.Lock() + self._record_lock = threading.Lock() + self._next_server_id = 9000 + self._pending: dict[Any, dict] = {} + self._pending_cv = threading.Condition() + self._initialized = False + self._interrupted: set[str] = set() + self._turn_thread: Optional[threading.Thread] = None + self.record({"event": "spawn", "argv": argv, "ppid": os.getppid()}) + + # --- io ----------------------------------------------------------------------------------------- + def record(self, entry: dict) -> None: + entry = {"pid": os.getpid(), "t": time.time(), **entry} + with self._record_lock, open(self.transcript_path, "a", encoding="utf-8") as fh: + fh.write(json.dumps(entry) + "\n") + + def send(self, msg: dict) -> None: + self.record({"dir": "out", "msg": msg}) + with self._write_lock: + sys.stdout.write(json.dumps(msg) + "\n") + sys.stdout.flush() + + def notify(self, method: str, params: dict) -> None: + self.send({"method": method, "params": params, "emittedAtMs": _now_ms()}) + + def error(self, rid: Any, message: str, code: int = INVALID_REQUEST) -> None: + self.send({"error": {"code": code, "message": message}, "id": rid}) + + # --- main loop ----------------------------------------------------------------------------------- + def serve(self) -> None: + for raw in sys.stdin: + raw = raw.strip() + if not raw: + continue + try: + msg = json.loads(raw) + except json.JSONDecodeError: + self.record({"dir": "in", "raw": raw, "violation": "not JSON"}) + continue + if "method" in msg: + self._on_client_message(msg) + else: + self._on_response(msg) + self.record({"event": "stdin_eof"}) + + def _on_client_message(self, msg: dict) -> None: + method, rid, params = msg.get("method"), msg.get("id"), msg.get("params") + if rid is None: # notification + known = method in {"initialized"} + self.record({"dir": "in", "msg": msg, **({} if known else {"violation": f"unknown notification {method}"})}) + return + spec = REQUEST_PARAMS.get(method) + if spec is None: + self.record({"dir": "in", "msg": msg, "violation": f"unknown method {method}"}) + self.error(rid, f"Invalid request: unknown variant `{method}`, expected one of " + + ", ".join(f"`{m}`" for m in REQUEST_PARAMS)) + return + err, ignored = validate(spec, params if params is not None else {}) + self.record({"dir": "in", "msg": msg, "ignored": ignored, **({"violation": err} if err else {})}) + if err: + self.error(rid, f"Invalid request: {err}") + return + if method != "initialize" and not self._initialized: + self.error(rid, "Not initialized") + return + self._HANDLERS[method](self, rid, params or {}) + + def _on_response(self, msg: dict) -> None: + with self._pending_cv: + pending = self._pending.get(msg.get("id")) + entry: dict = {"dir": "in", "msg": msg} + if pending is None: + entry["violation"] = "response to unknown server request id" + elif "result" in msg: + err, ignored = validate(SERVER_REQUEST_RESULTS[pending["method"]], msg["result"]) + entry.update({"reply_to": pending["method"], "ignored": ignored, **({"violation": err} if err else {})}) + else: + entry["reply_to"] = pending["method"] + self.record(entry) + if pending is not None: + with self._pending_cv: + pending["reply"] = msg + self._pending_cv.notify_all() + + def server_request(self, method: str, params: dict, timeout: float = 60.0) -> dict: + """Issue a server-initiated request and block for Hermes' reply.""" + with self._pending_cv: + self._next_server_id += 1 + rid = self._next_server_id + self._pending[rid] = {"method": method} + self.send({"id": rid, "method": method, "params": params}) + deadline = time.monotonic() + timeout + with self._pending_cv: + while "reply" not in self._pending[rid]: + remaining = deadline - time.monotonic() + if remaining <= 0: + self.record({"event": "server_request_unanswered", "method": method, "id": rid}) + return {} + self._pending_cv.wait(remaining) + return self._pending.pop(rid)["reply"] + + # --- request handlers ------------------------------------------------------------------------------ + def _initialize(self, rid: Any, params: dict) -> None: + if self._initialized: + self.error(rid, "Already initialized") + return + self._initialized = True + self.send({"id": rid, "result": { + "userAgent": f"{params['clientInfo']['name']}/{CLI_VERSION} (fake)", "codexHome": str(self.state_dir), + "platformFamily": "unix", "platformOs": sys.platform}}) + + def _thread_payload(self, thread_id: str, data: dict) -> dict: + info = data["threads"][thread_id] + return { + "id": thread_id, "sessionId": thread_id, "preview": info.get("preview", ""), "ephemeral": False, + "modelProvider": "openai", "createdAt": info["createdAt"], "updatedAt": int(time.time()), + "status": {"type": "idle"}, "cwd": info["cwd"], "cliVersion": CLI_VERSION, "source": "vscode", + "turns": [{"id": t["id"], "items": t["items"], "status": t["status"]} for t in info["turns"]], + } + + def _thread_response(self, thread_id: str, data: dict) -> dict: + info = data["threads"][thread_id] + return {"thread": self._thread_payload(thread_id, data), "model": "gpt-fake", "modelProvider": "openai", + "cwd": info["cwd"], "approvalPolicy": "on-request", "approvalsReviewer": "user", + "sandbox": {"type": "workspaceWrite"}, "instructionSources": []} + + def _thread_start(self, rid: Any, params: dict) -> None: + data = self.store.load() + thread_id = str(uuid.uuid4()) + data["threads"][thread_id] = {"createdAt": int(time.time()), "cwd": params.get("cwd") or os.getcwd(), + "turns": [], "developerInstructions": params.get("developerInstructions")} + self.store.save(data) + self.thread_id = thread_id + self.send({"id": rid, "result": self._thread_response(thread_id, data)}) + self.notify("thread/started", {"thread": self._thread_payload(thread_id, data)}) + + def _thread_resume(self, rid: Any, params: dict) -> None: + data = self.store.load() + thread_id = params["threadId"] + if self.scenario.get("forget_threads") or thread_id not in data["threads"]: + self.error(rid, f"no rollout found for thread id {thread_id}") + return + self.thread_id = thread_id + self.send({"id": rid, "result": self._thread_response(thread_id, data)}) + + def _turn_start(self, rid: Any, params: dict) -> None: + data = self.store.load() + thread_id = params["threadId"] + if thread_id not in data["threads"]: + self.error(rid, f"thread not found: {thread_id}") + return + turns = self.scenario.get("turns") or [] + cursor = data.get("turn_cursor", 0) + script = turns[cursor] if cursor < len(turns) else {"steps": [{"kind": "message", "text": "(unscripted)"}]} + data["turn_cursor"] = cursor + 1 + self.store.save(data) + if script.get("start_error"): + self.error(rid, script["start_error"]) + return + turn_id = str(uuid.uuid4()) + self.send({"id": rid, "result": {"turn": {"id": turn_id, "items": [], "status": "inProgress"}}}) + self._turn_thread = threading.Thread(target=self._play_turn, args=(thread_id, turn_id, params, script), + daemon=True) + self._turn_thread.start() + + def _turn_interrupt(self, rid: Any, params: dict) -> None: + self._interrupted.add(params["turnId"]) + self.send({"id": rid, "result": {}}) + + def _turn_steer(self, rid: Any, params: dict) -> None: + self.send({"id": rid, "result": {"turnId": params["expectedTurnId"]}}) + + def _compact_start(self, rid: Any, params: dict) -> None: + self.send({"id": rid, "result": {}}) + turn_id = str(uuid.uuid4()) + self._turn_thread = threading.Thread( + target=self._play_turn, args=(params["threadId"], turn_id, None, {"steps": [{"kind": "compaction"}]}), + daemon=True) + self._turn_thread.start() + + _HANDLERS: dict[str, Callable[..., None]] = { + "initialize": _initialize, "thread/start": _thread_start, "thread/resume": _thread_resume, + "turn/start": _turn_start, "turn/interrupt": _turn_interrupt, "turn/steer": _turn_steer, + "thread/compact/start": _compact_start, + } + + # --- turn playback ----------------------------------------------------------------------------------- + def _play_turn(self, thread_id: str, turn_id: str, params: Optional[dict], script: dict) -> None: + ctx = _TurnCtx(self, thread_id, turn_id) + self.notify("turn/started", {"threadId": thread_id, "turn": {"id": turn_id, "items": [], + "status": "inProgress"}}) + if params is not None: # codex echoes the submitted input as a userMessage item + ctx.item({"type": "userMessage", "id": ctx.new_id(), "content": params["input"]}) + status, error = "completed", None + try: + for step in script.get("steps", []): + ctx.check_interrupt() + STEPS[step["kind"]](ctx, step) + except _TurnEnded as ended: + status, error = ended.args[0], (ended.args[1] if len(ended.args) > 1 else None) + data = self.store.load() + data["threads"][thread_id]["turns"].append({"id": turn_id, "items": ctx.items, "status": status}) + self.store.save(data) + turn: dict = {"id": turn_id, "items": [], "status": status} + if error: + turn["error"] = {"message": error} + self.notify("turn/completed", {"threadId": thread_id, "turn": turn}) + + +class _TurnCtx: + def __init__(self, server: FakeAppServer, thread_id: str, turn_id: str) -> None: + self.server, self.thread_id, self.turn_id = server, thread_id, turn_id + self.items: list[dict] = [] + self._n = 0 + + def new_id(self) -> str: + self._n += 1 + return f"item_{self.turn_id[:8]}_{self._n}" + + def scope(self, **extra: Any) -> dict: + return {"threadId": self.thread_id, "turnId": self.turn_id, **extra} + + def started(self, item: dict) -> None: + self.server.notify("item/started", self.scope(item=item, startedAtMs=_now_ms())) + + def completed(self, item: dict) -> None: + self.items.append(item) + self.server.notify("item/completed", self.scope(item=item, completedAtMs=_now_ms())) + + def item(self, item: dict) -> None: + self.started(item) + self.completed(item) + + def check_interrupt(self) -> None: + if self.turn_id in self.server._interrupted: + raise _TurnEnded("interrupted") + + +def _step_reasoning(ctx: _TurnCtx, step: dict) -> None: + item_id = ctx.new_id() + ctx.started({"type": "reasoning", "id": item_id, "summary": [], "content": []}) + for index, part in enumerate(step.get("summary", [])): + for chunk in (part[: len(part) // 2], part[len(part) // 2:]): + ctx.server.notify("item/reasoning/summaryTextDelta", ctx.scope(itemId=item_id, delta=chunk, + summaryIndex=index)) + ctx.completed({"type": "reasoning", "id": item_id, "summary": step.get("summary", []), + "content": step.get("content", [])}) + + +def _step_command(ctx: _TurnCtx, step: dict) -> None: + item_id = ctx.new_id() + cwd = step.get("cwd") or ctx.server.store.load()["threads"][ctx.thread_id]["cwd"] + base = {"type": "commandExecution", "id": item_id, "command": step["command"], "cwd": cwd, + "commandActions": [{"type": "unknown", "command": step["command"]}]} + ctx.started({**base, "status": "inProgress"}) + decision = "accept" + if step.get("approval", True): + reply = ctx.server.server_request("item/commandExecution/requestApproval", ctx.scope( + itemId=item_id, startedAtMs=_now_ms(), command=step["command"], cwd=cwd, reason=step.get("reason"))) + decision = (reply.get("result") or {}).get("decision", "decline") + if decision not in ("accept", "acceptForSession"): + ctx.completed({**base, "status": "declined", "aggregatedOutput": None, "exitCode": None}) + return + ctx.server.notify("item/commandExecution/outputDelta", ctx.scope(itemId=item_id, delta=step["output"])) + ctx.completed({**base, "status": "completed" if step.get("exit_code", 0) == 0 else "failed", + "aggregatedOutput": step["output"], "exitCode": step.get("exit_code", 0), "durationMs": 7}) + + +def _step_message(ctx: _TurnCtx, step: dict) -> None: + item_id, text = ctx.new_id(), step["text"] + ctx.started({"type": "agentMessage", "id": item_id, "text": ""}) + size = max(1, len(text) // max(1, step.get("chunks", 3))) + for start in range(0, len(text), size): + ctx.server.notify("item/agentMessage/delta", ctx.scope(itemId=item_id, delta=text[start:start + size])) + ctx.completed({"type": "agentMessage", "id": item_id, "text": text}) + + +def _step_message_partial(ctx: _TurnCtx, step: dict) -> None: + """An agentMessage that streams deltas but never reaches item/completed (dies mid-item).""" + item_id = ctx.new_id() + ctx.started({"type": "agentMessage", "id": item_id, "text": ""}) + ctx.server.notify("item/agentMessage/delta", ctx.scope(itemId=item_id, delta=step["text"])) + + +def _step_usage(ctx: _TurnCtx, step: dict) -> None: + breakdown = {"inputTokens": step["input"], "cachedInputTokens": step.get("cached", 0), + "outputTokens": step.get("output", 10), "reasoningOutputTokens": step.get("reasoning", 0)} + breakdown["totalTokens"] = breakdown["inputTokens"] + breakdown["outputTokens"] + ctx.server.notify("thread/tokenUsage/updated", ctx.scope(tokenUsage={ + "last": breakdown, "total": breakdown, "modelContextWindow": step.get("window", 272000)})) + + +def _step_compaction(ctx: _TurnCtx, step: dict) -> None: + ctx.item({"type": "contextCompaction", "id": ctx.new_id()}) + + +def _step_error_note(ctx: _TurnCtx, step: dict) -> None: + ctx.server.notify("error", ctx.scope(error={"message": step["message"]}, willRetry=step.get("will_retry", True))) + + +def _step_fail(ctx: _TurnCtx, step: dict) -> None: + raise _TurnEnded("failed", step["message"]) + + +def _step_crash(ctx: _TurnCtx, step: dict) -> None: + ctx.server.record({"event": "crash", "code": step.get("code", 1)}) + sys.stderr.write(step.get("stderr", "fake codex: fatal error") + "\n") + sys.stderr.flush() + os._exit(step.get("code", 1)) + + +def _step_progress(ctx: _TurnCtx, step: dict) -> None: + """Keep the turn busy for ``seconds``, streaming a delta every ``interval`` (a model that is working).""" + item_id = ctx.new_id() + ctx.started({"type": "agentMessage", "id": item_id, "text": ""}) + deadline = time.monotonic() + step["seconds"] + text = "" + while time.monotonic() < deadline: + ctx.check_interrupt() + text += "." + ctx.server.notify("item/agentMessage/delta", ctx.scope(itemId=item_id, delta=".")) + time.sleep(step.get("interval", 0.25)) + ctx.check_interrupt() + ctx.completed({"type": "agentMessage", "id": item_id, "text": step.get("text", text)}) + + +def _step_permissions(ctx: _TurnCtx, step: dict) -> None: + ctx.server.server_request("item/permissions/requestApproval", ctx.scope( + itemId=ctx.new_id(), startedAtMs=_now_ms(), cwd=os.getcwd(), reason=step.get("reason"), + permissions={"network": {"enabled": True}})) + + +# The descendant lives until the harness drops the release file (or 600 s pass), so the test side can +# retire an orphan without signalling a process that is no longer in its own subtree. +_GRANDCHILD = ("import os, sys, time\n" + "deadline = time.monotonic() + 600\n" + "while time.monotonic() < deadline and not os.path.exists(sys.argv[1]):\n" + " time.sleep(0.2)\n") + + +def _step_grandchild(ctx: _TurnCtx, step: dict) -> None: + """A descendant in its own session (like codex's stdio MCP servers).""" + child = subprocess.Popen([sys.executable, "-c", _GRANDCHILD, str(ctx.server.state_dir / GRANDCHILD_RELEASE)], + stdin=subprocess.DEVNULL, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, + start_new_session=True) + ctx.server.record({"event": "grandchild", "child_pid": child.pid}) + + +STEPS: dict[str, Callable[[_TurnCtx, dict], None]] = { + "reasoning": _step_reasoning, "command": _step_command, "message": _step_message, "usage": _step_usage, + "message_partial": _step_message_partial, + "compaction": _step_compaction, "error_note": _step_error_note, "fail": _step_fail, "crash": _step_crash, + "progress": _step_progress, "permissions": _step_permissions, "grandchild": _step_grandchild, +} + + +def main(argv: list[str]) -> int: + if "--version" in argv: + print(f"codex-cli {CLI_VERSION}") + return 0 + state_dir = Path(argv[argv.index("--state-dir") + 1]) + if "app-server" not in argv: + sys.stderr.write(f"fake codex: unsupported argv {argv}\n") + return 2 + server = FakeAppServer(state_dir, argv) + + def on_sigterm(*_: Any) -> None: + server.record({"event": "exit", "how": "SIGTERM"}) + os._exit(143) + + signal.signal(signal.SIGTERM, on_sigterm) + server.serve() + server.record({"event": "exit", "how": "stdin_eof"}) + return 0 + + +# --------------------------------------------------------------------------------------------------- +# Test-side harness +# --------------------------------------------------------------------------------------------------- + +class FakeCodex: + """One fake codex install: wrapper executable + state dir, and readers over the transcript.""" + + def __init__(self, root: Path, turns: list[dict], **scenario: Any) -> None: + self.state_dir = root / "fake_codex" + self.state_dir.mkdir(parents=True, exist_ok=True) + (self.state_dir / "scenario.json").write_text(json.dumps({"turns": turns, **scenario}), encoding="utf-8") + self.bin = root / "bin" / "codex" + self.bin.parent.mkdir(parents=True, exist_ok=True) + self.bin.write_text(f'#!/bin/sh\nexec "{sys.executable}" "{Path(__file__).resolve()}" ' + f'--state-dir "{self.state_dir}" "$@"\n', encoding="utf-8") + self.bin.chmod(0o755) + + def set_scenario(self, **changes: Any) -> None: + path = self.state_dir / "scenario.json" + data = json.loads(path.read_text(encoding="utf-8")) + data.update(changes) + path.write_text(json.dumps(data), encoding="utf-8") + + def entries(self) -> list[dict]: + path = self.state_dir / "transcript.jsonl" + if not path.exists(): + return [] + return [json.loads(line) for line in path.read_text(encoding="utf-8").splitlines() if line.strip()] + + def requests(self, method: Optional[str] = None) -> list[dict]: + """Client requests Hermes sent (full transcript entries), optionally one method.""" + return [e for e in self.entries() if e.get("dir") == "in" and "method" in e.get("msg", {}) + and "id" in e["msg"] and (method is None or e["msg"]["method"] == method)] + + def replies_to(self, server_method: str) -> list[dict]: + return [e for e in self.entries() if e.get("reply_to") == server_method] + + def violations(self) -> list[dict]: + return [e for e in self.entries() if e.get("violation")] + + def ignored_fields(self) -> list[tuple[str, str]]: + return [(str(e["msg"].get("method") or e.get("reply_to")), path) for e in self.entries() + for path in e.get("ignored") or []] + + def spawned_pids(self) -> list[int]: + return [e["pid"] for e in self.entries() if e.get("event") == "spawn"] + + def grandchild_pids(self) -> list[int]: + return [e["child_pid"] for e in self.entries() if e.get("event") == "grandchild"] + + def threads(self) -> dict: + return _Store(self.state_dir / "threads.json").load()["threads"] + + def assert_wire_clean(self) -> None: + """Every request/response Hermes sent is valid AND carries no field codex would silently drop.""" + assert not self.violations(), f"protocol violations: {self.violations()}" + assert not self.ignored_fields(), f"fields codex ignores (intent silently lost): {self.ignored_fields()}" + + +@dataclass +class CodexRun: + """Outcome of :func:`run_codex_scenario`: the fake, the Hermes home and one ChatResult per CLI run.""" + fake: FakeCodex + home: Any + results: list + session_id: str + + @property + def output(self) -> str: + return "".join(r.stdout + r.stderr for r in self.results) + + def process_entries(self, index: int) -> list[dict]: + """Transcript entries of the ``index``-th app-server process (one per CLI run).""" + pid = self.fake.spawned_pids()[index] + return [e for e in self.fake.entries() if e.get("pid") == pid] + + def process_requests(self, index: int, method: str) -> list[dict]: + return [e["msg"] for e in self.process_entries(index) + if e.get("dir") == "in" and e.get("msg", {}).get("method") == method and "id" in e["msg"]] + + def cleanup(self) -> None: + """Retire anything the fake spawned that is still alive (orphan scenarios). + + Orphaned descendants are released cooperatively first: by teardown they are reparented to init, + outside the test's process subtree, where a live-system guard (rightly) refuses to signal them.""" + (self.fake.state_dir / GRANDCHILD_RELEASE).touch() + deadline = time.monotonic() + 10 + while time.monotonic() < deadline and any(pid_alive(p) for p in self.fake.grandchild_pids()): + time.sleep(0.05) + for pid in self.fake.spawned_pids() + self.fake.grandchild_pids(): + if pid_alive(pid): + with contextlib.suppress(OSError): + os.kill(pid, signal.SIGKILL) + + +def run_codex_scenario(root: Path, turns: list[dict], runs: list[dict], *, config: Optional[dict] = None, + **scenario: Any) -> CodexRun: + """Real ``hermes chat -q`` runs (``--resume`` after the first) against a fresh fake codex install. + + ``runs``: ``{"prompt", "args": [...], "then": {scenario changes applied after this run}}``.""" + from tests.e2e.core.providers._native_helpers import latest_session, make_home, run_chat + + fake = FakeCodex(root, turns, **scenario) + model = {"provider": "openai", "default": "gpt-5.5", "openai_runtime": "codex_app_server", + "codex_bin": str(fake.bin)} + # The app-server owns auth; the key only satisfies Hermes' provider resolution and never leaves. + home = make_home(root, model, env_file={"OPENAI_API_KEY": "sk-fake-codex-e2e"}, extra_config=config) + results, session_id = [], None + for run in runs: + results.append(run_chat(home, run["prompt"], args=tuple(run.get("args", ())), resume=session_id, + timeout=run.get("timeout", 120))) + session_id = latest_session(home) + if run.get("then"): + fake.set_scenario(**run["then"]) + assert session_id is not None + return CodexRun(fake, home, results, session_id) + + +def pid_alive(pid: int) -> bool: + """True while ``pid`` exists and is not a zombie.""" + try: + with open(f"/proc/{pid}/stat", encoding="utf-8") as fh: + return fh.read().rsplit(")", 1)[1].split()[0] != "Z" + except (FileNotFoundError, IndexError, ProcessLookupError): + return False + + +if __name__ == "__main__": + sys.exit(main(sys.argv[1:])) diff --git a/tests/fakes/providers/copilot_acp.py b/tests/fakes/providers/copilot_acp.py new file mode 100644 index 000000000000..9fd9bb336945 --- /dev/null +++ b/tests/fakes/providers/copilot_acp.py @@ -0,0 +1,465 @@ +#!/usr/bin/env python3 +"""Fake ACP agent executable standing in for ``copilot --acp --stdio`` (provider ``copilot-acp``). + +Hermes' ``copilot-acp`` provider spawns an external agent process per model call and speaks the +Agent Client Protocol to it: JSON-RPC 2.0, one JSON object per line over stdio +(https://agentclientprotocol.com/protocol/overview). This module is that process. It + +* answers `` --help`` with a usage text advertising ``--acp`` (Hermes probes it before spawning); +* validates every client request against the published ACP schema (the ``agent-client-protocol`` + package's pydantic models, ``acp.schema``) plus the spec rules the models do not encode + (initialize-first, absolute ``cwd``, no custom root fields, known ``sessionId``) and REJECTS + malformed ones with a JSON-RPC ``-32602 Invalid params`` / ``-32600`` error, as an agent would; +* builds every message it sends from the same schema models, so a shape drift fails here first; +* replays a scripted turn per main-turn ``session/prompt`` (thought / message chunks with delays, + ``tool_call`` / ``tool_call_update``, ``session/request_permission``, ``fs/read_text_file``, + the ``{stopReason}`` result, JSON-RPC errors, a hard crash, chunks AFTER the result); +* appends every inbound and outbound message, with wall-clock time and pid, to + ``/transcript.jsonl`` (shared by every process of one scenario). + +Test-side API: :class:`AcpFake` (write the script, build the launcher, read the transcript) and the +action builders (:func:`thought`, :func:`message`, ...). Run as a script it is the agent itself: +``copilot_acp.py --acp --stdio --state ``. +""" + +from __future__ import annotations + +import fcntl +import json +import os +import signal +import sys +import time +from pathlib import Path +from typing import Any + +from pydantic import ValidationError + +PROTOCOL_VERSION = 1 +TOOLS_MARKER = "Available tools" +AUX_ANSWER = "fake-acp auxiliary answer" +USAGE = """usage: copilot [options] + +Options: + --acp Run as an Agent Client Protocol server + --stdio Use stdio transport (with --acp) + --state (fake) scenario directory holding script.json / transcript.jsonl + -h, --help Show help +""" + + +# ── test-side API ──────────────────────────────────────────────────────────────────────────── + + +def thought(text: str, delay: float = 0.0) -> dict[str, Any]: + return {"type": "thought", "text": text, "delay": delay} + + +def message(text: str, delay: float = 0.0) -> dict[str, Any]: + return {"type": "message", "text": text, "delay": delay} + + +def tool_call(call_id: str, title: str, kind: str = "other", status: str = "pending") -> dict[str, Any]: + return {"type": "tool_call", "id": call_id, "title": title, "kind": kind, "status": status} + + +def tool_update(call_id: str, status: str = "completed", text: str = "") -> dict[str, Any]: + return {"type": "tool_update", "id": call_id, "status": status, "text": text} + + +def permission(call_id: str, title: str = "run a command") -> dict[str, Any]: + return {"type": "permission", "id": call_id, "title": title} + + +def fs_read(path: str) -> dict[str, Any]: + return {"type": "fs_read", "path": path} + + +def result(stop_reason: str = "end_turn") -> dict[str, Any]: + return {"type": "result", "stopReason": stop_reason} + + +def rpc_error(code: int, msg: str, data: Any = None) -> dict[str, Any]: + return {"type": "error", "code": code, "message": msg, "data": data} + + +def crash(exit_code: int = 1, stderr: str = "fatal: agent crashed") -> dict[str, Any]: + return {"type": "crash", "code": exit_code, "stderr": stderr} + + +def hermes_tool_call(call_id: str, name: str, args: dict[str, Any]) -> str: + """Text a model behind ACP emits to call a Hermes tool (ACP has no OpenAI tools channel).""" + body = {"id": call_id, "type": "function", "function": {"name": name, "arguments": json.dumps(args)}} + return f"{json.dumps(body)}" + + +class AcpFake: + """One scenario directory: ``script.json`` in, ``transcript.jsonl`` out, plus a launcher.""" + + def __init__(self, state_dir: Path, turns: list[list[dict[str, Any]]], *, models: list[str] | None = None, + current_model: str | None = None, ignore_sigterm: bool = False, load_session: bool = True, + aux_text: str = AUX_ANSWER): + self.state_dir = Path(state_dir) + self.state_dir.mkdir(parents=True, exist_ok=True) + script = {"turns": turns, "models": models or [], "current_model": current_model, + "ignore_sigterm": ignore_sigterm, "load_session": load_session, "aux_text": aux_text} + (self.state_dir / "script.json").write_text(json.dumps(script), encoding="utf-8") + self.launcher = self.state_dir / "copilot" + self.launcher.write_text( + f'#!/bin/sh\nexec "{sys.executable}" "{Path(__file__).resolve()}" "$@"\n', encoding="utf-8") + self.launcher.chmod(0o755) + + def env(self) -> dict[str, str]: + """Env vars that point Hermes' copilot-acp client at this fake.""" + return {"HERMES_COPILOT_ACP_COMMAND": str(self.launcher), + "HERMES_COPILOT_ACP_ARGS": f"--acp --stdio --state {self.state_dir}"} + + def records(self) -> list[dict[str, Any]]: + path = self.state_dir / "transcript.jsonl" + if not path.exists(): + return [] + return [json.loads(line) for line in path.read_text(encoding="utf-8").splitlines() if line.strip()] + + def inbound(self, method: str | None = None) -> list[dict[str, Any]]: + """Client->agent messages (requests, notifications, responses), optionally one method.""" + return [r for r in self.records() if r["dir"] == "in" and (method is None or r["msg"].get("method") == method)] + + def main_prompts(self) -> list[dict[str, Any]]: + """Main-turn ``session/prompt`` records (those carrying Hermes' tool bridge).""" + return [r for r in self.inbound("session/prompt") if r.get("main")] + + def aux_prompts(self) -> list[dict[str, Any]]: + """Auxiliary ``session/prompt`` records (no tool bridge: compaction summaries, titles).""" + return [r for r in self.inbound("session/prompt") if not r.get("main")] + + def invalid(self) -> list[dict[str, Any]]: + return [r for r in self.records() if r["dir"] == "in" and r.get("errors")] + + def pids(self) -> list[int]: + return sorted({r["pid"] for r in self.records() if r["dir"] == "spawn"}) + + def events(self, kind: str) -> list[dict[str, Any]]: + return [r for r in self.records() if r["dir"] == kind] + + +def prompt_text(record: dict[str, Any]) -> str: + return "".join(b.get("text", "") for b in record["msg"]["params"].get("prompt", []) if b.get("type") == "text") + + +# ── agent process ──────────────────────────────────────────────────────────────────────────── + + +class Agent: + """One spawned agent process (one Hermes model call).""" + + def __init__(self, state_dir: Path): + import acp.schema as schema # validation/building uses the published ACP models + + self.s = schema + self.state_dir = state_dir + self.script = json.loads((state_dir / "script.json").read_text(encoding="utf-8")) + self.initialized = False + self.client_caps: dict[str, Any] = {} + self.sessions: set[str] = set() + self.turn: int | None = None + self.next_id = 1000 + self.request_schemas = { + "initialize": schema.InitializeRequest, + "session/new": schema.NewSessionRequest, + "session/load": schema.LoadSessionRequest, + "session/prompt": schema.PromptRequest, + "session/set_model": schema.SetSessionModelRequest, + "session/set_config_option": schema.SetSessionConfigOptionSelectRequest, + "session/cancel": schema.CancelNotification, + } + self.handlers = { + "initialize": self._initialize, "session/new": self._new_session, "session/load": self._load_session, + "session/set_model": self._ack, "session/set_config_option": self._set_config_option, + "session/prompt": self._prompt, + } + + # transcript + wire ------------------------------------------------------------------------ + + def log(self, direction: str, msg: Any = None, **extra: Any) -> None: + rec = {"t": time.time(), "pid": os.getpid(), "turn": self.turn, "dir": direction, "msg": msg, **extra} + with open(self.state_dir / "transcript.jsonl", "a", encoding="utf-8") as fh: + fh.write(json.dumps(rec) + "\n") + + def send(self, msg: dict[str, Any]) -> None: + self.log("out", msg) + sys.stdout.write(json.dumps(msg) + "\n") + sys.stdout.flush() + + def dump(self, model: Any) -> dict[str, Any]: + return model.model_dump(by_alias=True, exclude_none=True, mode="json") + + def reply(self, msg_id: Any, result_model: Any) -> None: + body = result_model if isinstance(result_model, dict) or result_model is None else self.dump(result_model) + self.send({"jsonrpc": "2.0", "id": msg_id, "result": body}) + + def error(self, msg_id: Any, code: int, text: str, data: Any = None) -> None: + err = {"code": code, "message": text, **({"data": data} if data is not None else {})} + self.send({"jsonrpc": "2.0", "id": msg_id, "error": err}) + + def update(self, session_id: str, update_model: Any) -> None: + note = self.s.SessionNotification(session_id=session_id, update=update_model) + self.send({"jsonrpc": "2.0", "method": "session/update", "params": self.dump(note)}) + + def read(self) -> dict[str, Any] | None: + line = sys.stdin.readline() + if not line: + return None + try: + msg = json.loads(line) + except json.JSONDecodeError as exc: + self.log("in", line.rstrip("\n"), errors=[f"parse error: {exc}"]) + self.error(None, -32700, "Parse error") + return {} + return msg + + # validation ------------------------------------------------------------------------------- + + def validate(self, msg: dict[str, Any]) -> list[str]: + """Spec violations of one client request/notification (empty list = valid).""" + errors: list[str] = [] + if msg.get("jsonrpc") != "2.0": + errors.append("jsonrpc must be '2.0'") + method = msg.get("method") + if "id" in msg and not isinstance(msg["id"], (int, str)): + errors.append("id must be a string or integer") + model = self.request_schemas.get(method) + if model is None: + return errors + params = msg.get("params") + try: + parsed = model.model_validate(params) + except ValidationError as exc: + return errors + [f"schema: {e['loc']} {e['msg']}" for e in exc.errors()] + allowed = {f.alias or name for name, f in model.model_fields.items()} + errors += [f"custom root field {k!r} (spec: use _meta)" for k in sorted(set(params) - allowed)] + if method != "initialize" and not self.initialized: + errors.append("request before initialize") + if method in ("session/new", "session/load") and not os.path.isabs(parsed.cwd): + errors.append("cwd must be an absolute path") + sid = getattr(parsed, "session_id", None) + if method not in ("session/load",) and sid is not None and sid not in self.sessions: + errors.append(f"unknown sessionId {sid!r}") + return errors + + # request handlers ------------------------------------------------------------------------- + + def _initialize(self, msg_id: Any, params: dict[str, Any]) -> None: + self.initialized = True + self.client_caps = params.get("clientCapabilities") or {} + caps = self.s.AgentCapabilities(load_session=bool(self.script.get("load_session"))) + self.reply(msg_id, self.s.InitializeResponse( + protocol_version=PROTOCOL_VERSION, agent_capabilities=caps, auth_methods=[], + agent_info=self.s.Implementation(name="fake-copilot", title="Fake Copilot", version="1.0.0"))) + + def _session_response(self, session_id: str) -> Any: + models = self.script.get("models") or [] + options = None + if models: + current = self.script.get("current_model") or models[0] + options = [self.s.SessionConfigOptionSelect( + id="model", name="Model", category="model", type="select", current_value=current, + options=[self.s.SessionConfigSelectOption(value=m, name=m) for m in models])] + return self.s.NewSessionResponse(session_id=session_id, config_options=options) + + def _new_session(self, msg_id: Any, params: dict[str, Any]) -> None: + session_id = f"fake-sess-{os.getpid()}-{len(self.sessions) + 1}" + self.sessions.add(session_id) + self.reply(msg_id, self._session_response(session_id)) + + def _load_session(self, msg_id: Any, params: dict[str, Any]) -> None: + self.sessions.add(params["sessionId"]) + self.reply(msg_id, self.s.LoadSessionResponse()) + + def _ack(self, msg_id: Any, params: dict[str, Any]) -> None: + self.reply(msg_id, {}) + + def _set_config_option(self, msg_id: Any, params: dict[str, Any]) -> None: + models = self.script.get("models") or [] + if params.get("configId") != "model" or params.get("value") not in models: + self.error(msg_id, -32602, "Invalid params", {"reason": "unknown config option or value"}) + return + self.script["current_model"] = params["value"] + self.reply(msg_id, self.s.SetSessionConfigOptionResponse(config_options=self._session_response("x").config_options)) + + def _claim_turn(self) -> int: + """Next scripted turn index, shared across every process of the scenario (file lock).""" + with open(self.state_dir / "turn.counter", "a+", encoding="utf-8") as fh: + fcntl.flock(fh, fcntl.LOCK_EX) + fh.seek(0) + index = int(fh.read().strip() or 0) + fh.seek(0) + fh.truncate() + fh.write(str(index + 1)) + return index + + def _prompt(self, msg_id: Any, params: dict[str, Any]) -> None: + text = "".join(b.get("text", "") for b in params["prompt"] if b.get("type") == "text") + session_id = params["sessionId"] + if TOOLS_MARKER not in text: # auxiliary call (title/summary): never consumes a scripted turn + self.update(session_id, self._chunk("agent_message_chunk", self.script.get("aux_text") or AUX_ANSWER)) + self.reply(msg_id, self.s.PromptResponse(stop_reason="end_turn")) + return + turns = self.script.get("turns") or [] + index = self._claim_turn() + self.turn = index + self.log("turn", {"index": index}) + actions = turns[index] if index < len(turns) else [message("(fake-acp: script exhausted)")] + replied = False + for action in actions: + replied = self._run_action(action, msg_id, session_id) or replied + if not replied: + self.reply(msg_id, self.s.PromptResponse(stop_reason="end_turn")) + + # scripted actions ------------------------------------------------------------------------- + + def _chunk(self, kind: str, text: str) -> Any: + cls = {"agent_message_chunk": self.s.AgentMessageChunk, "agent_thought_chunk": self.s.AgentThoughtChunk}[kind] + return cls(session_update=kind, content=self.s.TextContentBlock(type="text", text=text)) + + def _run_action(self, action: dict[str, Any], msg_id: Any, session_id: str) -> bool: + """Perform one scripted action; True when it answered the pending ``session/prompt``.""" + if action.get("delay"): + time.sleep(float(action["delay"])) + return bool(self._actions[action["type"]](self, action, msg_id, session_id)) + + def _a_thought(self, action, msg_id, session_id): + self.update(session_id, self._chunk("agent_thought_chunk", action["text"])) + + def _a_message(self, action, msg_id, session_id): + self.update(session_id, self._chunk("agent_message_chunk", action["text"])) + + def _a_tool_call(self, action, msg_id, session_id): + self.update(session_id, self.s.ToolCallStart( + session_update="tool_call", tool_call_id=action["id"], title=action["title"], kind=action["kind"], + status=action["status"])) + + def _a_tool_update(self, action, msg_id, session_id): + content = None + if action.get("text"): + content = [self.s.ContentToolCallContent( + type="content", content=self.s.TextContentBlock(type="text", text=action["text"]))] + self.update(session_id, self.s.ToolCallProgress( + session_update="tool_call_update", tool_call_id=action["id"], status=action["status"], content=content)) + + def _a_permission(self, action, msg_id, session_id): + req = self.s.RequestPermissionRequest( + session_id=session_id, + tool_call=self.s.ToolCallUpdate(tool_call_id=action["id"], title=action["title"], kind="execute"), + options=[self.s.PermissionOption(option_id="allow-once", name="Allow once", kind="allow_once"), + self.s.PermissionOption(option_id="reject-once", name="Reject", kind="reject_once")]) + response = self._client_request("session/request_permission", self.dump(req)) + self._check_response(response, self.s.RequestPermissionResponse, "permission_outcome") + + def _a_fs_read(self, action, msg_id, session_id): + if not (self.client_caps.get("fs") or {}).get("readTextFile"): + self.log("fs_skipped", {"reason": "client did not advertise fs.readTextFile"}) + return + req = self.s.ReadTextFileRequest(session_id=session_id, path=action["path"]) + response = self._client_request("fs/read_text_file", self.dump(req)) + self._check_response(response, self.s.ReadTextFileResponse, "fs_read_result") + + def _a_result(self, action, msg_id, session_id): + self.reply(msg_id, self.s.PromptResponse(stop_reason=action["stopReason"])) + return True + + def _a_error(self, action, msg_id, session_id): + self.error(msg_id, int(action["code"]), action["message"], action.get("data")) + return True + + def _a_crash(self, action, msg_id, session_id): + self.log("crash", {"code": action["code"]}) + sys.stderr.write(action["stderr"] + "\n") + sys.stderr.flush() + os._exit(int(action["code"])) + + _actions = {"thought": _a_thought, "message": _a_message, "tool_call": _a_tool_call, "tool_update": _a_tool_update, + "permission": _a_permission, "fs_read": _a_fs_read, "result": _a_result, "error": _a_error, + "crash": _a_crash} + + def _client_request(self, method: str, params: dict[str, Any]) -> dict[str, Any] | None: + self.next_id += 1 + req_id = self.next_id + self.send({"jsonrpc": "2.0", "id": req_id, "method": method, "params": params}) + while (msg := self.read()) is not None: + if msg.get("id") == req_id and "method" not in msg: + return msg + self.dispatch(msg) + return None + + def _check_response(self, response: dict[str, Any] | None, model: Any, kind: str) -> None: + errors: list[str] = [] + if response is None: + errors.append("client closed stdin before answering") + elif "result" in response: + try: + model.model_validate(response["result"]) + except ValidationError as exc: + errors += [f"schema: {e['loc']} {e['msg']}" for e in exc.errors()] + elif not isinstance((response or {}).get("error"), dict): + errors.append("response carries neither result nor error") + self.log("in", response, kind=kind, errors=errors) + + # loop ------------------------------------------------------------------------------------- + + def dispatch(self, msg: dict[str, Any]) -> None: + method = msg.get("method") + errors = self.validate(msg) if method else ["unsolicited response"] + main = method == "session/prompt" and TOOLS_MARKER in "".join( + b.get("text", "") for b in ((msg.get("params") or {}).get("prompt") or []) if isinstance(b, dict)) + self.log("in", msg, errors=errors, main=main) + if "id" not in msg or not method: + return # notification (session/cancel) or stray response + if errors: + code = -32600 if errors == ["request before initialize"] else -32602 + self.error(msg["id"], code, "Invalid params" if code == -32602 else "Invalid request", {"errors": errors}) + return + handler = self.handlers.get(method) + if handler is None: + self.error(msg["id"], -32601, f"Method not found: {method}") + return + handler(msg["id"], msg.get("params") or {}) + + def serve(self) -> int: + self.log("spawn", {"argv": sys.argv[1:], "cwd": os.getcwd()}) + while (msg := self.read()) is not None: + if msg: + self.dispatch(msg) + if self.script.get("ignore_sigterm"): + # A wedged agent: ignores SIGTERM AND stdin EOF; only SIGKILL ends it. + self.log("wedged", {"reason": "stdin closed; ignoring EOF"}) + while True: + time.sleep(3600) + self.log("exit", {"reason": "stdin closed"}) + return 0 + + +def _install_signals(agent: Agent) -> None: + def _on_term(signum, _frame): + if agent.script.get("ignore_sigterm"): + agent.log("signal", {"signal": signum, "ignored": True}) + return + agent.log("exit", {"reason": f"signal {signum}"}) + os._exit(0) + + signal.signal(signal.SIGTERM, _on_term) + + +def main(argv: list[str]) -> int: + if not argv or "-h" in argv or "--help" in argv: + sys.stdout.write(USAGE) + return 0 + if "--acp" not in argv or "--state" not in argv: + sys.stderr.write("error: unknown option; this fake only runs with --acp --stdio --state \n") + return 1 + agent = Agent(Path(argv[argv.index("--state") + 1])) + _install_signals(agent) + return agent.serve() + + +if __name__ == "__main__": + sys.exit(main(sys.argv[1:])) diff --git a/tests/fakes/providers/gemini_native.py b/tests/fakes/providers/gemini_native.py new file mode 100644 index 000000000000..33c8575630c2 --- /dev/null +++ b/tests/fakes/providers/gemini_native.py @@ -0,0 +1,728 @@ +"""Fake Google AI Studio ``generateContent`` / ``streamGenerateContent`` endpoint behind a real +TLS boundary. + +Hermes routes to its native Gemini adapter only for the real Google host +(``generativelanguage.googleapis.com``), so the fake is reached the way any corporate egress +proxy would be: an HTTPS ``CONNECT`` proxy on loopback that terminates TLS for the Google host +with a leaf certificate signed by a throwaway CA. The child trusts that CA through the standard +``SSL_CERT_FILE`` / ``REQUESTS_CA_BUNDLE`` channel and reaches the proxy through +``HTTPS_PROXY`` (see :meth:`GeminiFake.child_env`). Every other host is refused and recorded, +so a test also proves the turn made no other egress. + +Requests are validated against the published Gemini API reference (``google.ai.generativelanguage`` +``v1beta`` / ``v1``: ``GenerateContentRequest``, ``Content``, ``Part``, ``Tool``, +``FunctionDeclaration``, ``Schema``) and rejected the way Google does: HTTP 400 with the +``{"error": {"code", "message", "status": "INVALID_ARGUMENT"}}`` body. That covers unknown proto +fields, role alternation, functionCall/functionResponse pairing (count, name, id), the Gemini 3 +thought-signature rule for the current turn, and the OpenAPI ``Schema`` subset of +``FunctionDeclaration.parameters``. Scripted replies are built from the vendor's +``GenerateContentResponse`` shape, including fault injection (Google error bodies, blocked +candidates, a mid-stream connection drop). +""" + +from __future__ import annotations + +import base64 +import datetime +import ipaddress +import json +import re +import secrets +import socket +import ssl +import struct +import threading +import time +from dataclasses import dataclass, field +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from pathlib import Path +from typing import Any, Callable +from urllib.parse import parse_qs, urlsplit + +GEMINI_HOST = "generativelanguage.googleapis.com" +MODEL_ID = "gemini-3-flash-preview" +API_KEY = "AIzaFakeGeminiKeyForHermesE2E0000000000" +# Documented dummy signatures that tell Gemini 3 to skip thought-signature validation. +SKIP_SIGNATURES = frozenset({"skip_thought_signature_validator", "context_engineering_is_the_way_to_go"}) +# Hermes-side config for a home that talks to this fake: the user-facing provider id + model and the +# key in ``.env`` (``model.base_url`` may pin another Google API version, e.g. ``.../v1``). +HERMES_ENV = {"GEMINI_API_KEY": API_KEY} + + +def hermes_model(base_url: str | None = None, **extra: Any) -> dict[str, Any]: + model: dict[str, Any] = {"provider": "gemini", "default": MODEL_ID, **extra} + if base_url: + model["base_url"] = base_url + return model + +# ── published proto field sets (JSON names; proto JSON parsing also accepts snake_case) ───────── +_REQUEST_FIELDS = {"contents", "tools", "toolConfig", "safetySettings", "systemInstruction", + "generationConfig", "cachedContent"} +_CONTENT_FIELDS = {"role", "parts"} +_PART_DATA_FIELDS = {"text", "inlineData", "functionCall", "functionResponse", "fileData", + "executableCode", "codeExecutionResult"} +_PART_FIELDS = _PART_DATA_FIELDS | {"thought", "thoughtSignature", "videoMetadata", "partMetadata"} +_FUNCTION_CALL_FIELDS = {"id", "name", "args"} +_FUNCTION_RESPONSE_FIELDS = {"id", "name", "response", "parts", "willContinue", "scheduling"} +_TOOL_FIELDS = {"functionDeclarations", "googleSearchRetrieval", "codeExecution", "googleSearch", + "urlContext", "computerUse", "fileSearch", "googleMaps"} +_DECL_FIELDS_V1BETA = {"name", "description", "behavior", "parameters", "parametersJsonSchema", + "response", "responseJsonSchema"} +_DECL_FIELDS_V1 = {"name", "description", "behavior", "parameters", "response"} +_GENERATION_FIELDS = {"stopSequences", "responseMimeType", "responseSchema", "responseJsonSchema", + "responseModalities", "candidateCount", "maxOutputTokens", "temperature", "topP", + "topK", "seed", "presencePenalty", "frequencyPenalty", "responseLogprobs", + "logprobs", "enableEnhancedCivicAnswers", "speechConfig", "thinkingConfig", + "mediaResolution", "imageConfig"} +_THINKING_FIELDS = {"includeThoughts", "thinkingBudget", "thinkingLevel"} +# google.ai.generativelanguage.v1beta.Schema (the OpenAPI 3.0 subset behind ``parameters``). +_SCHEMA_FIELDS = {"type", "format", "title", "description", "nullable", "enum", "maxItems", "minItems", + "properties", "required", "minProperties", "maxProperties", "minLength", "maxLength", + "pattern", "example", "anyOf", "propertyOrdering", "default", "items", "minimum", + "maximum"} +_SCHEMA_TYPES = {"TYPE_UNSPECIFIED", "STRING", "NUMBER", "INTEGER", "BOOLEAN", "ARRAY", "OBJECT", "NULL"} +_FUNCTION_NAME_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_.:\-]{0,127}$") +_ROLES = {"user", "model"} +_PATH_RE = re.compile(r"^/(?Pv1(?:beta|alpha)?)/models/(?P[^/:]+):(?P[A-Za-z]+)$") + + +class InvalidArgument(Exception): + """A request Google would refuse with HTTP 400 INVALID_ARGUMENT.""" + + +def _camel(key: str) -> str: + head, *rest = key.split("_") + return head + "".join(p[:1].upper() + p[1:] for p in rest) + + +def _check_fields(obj: Any, allowed: set[str], where: str) -> dict[str, Any]: + if not isinstance(obj, dict): + raise InvalidArgument(f"Invalid value at '{where}': expected an object") + for key in obj: + if _camel(key) not in allowed: + raise InvalidArgument(f"Invalid JSON payload received. Unknown name \"{key}\" at '{where}': " + "Cannot find field.") + return {_camel(k): v for k, v in obj.items()} + + +# ── scripted replies ──────────────────────────────────────────────────────────────────────────── +@dataclass +class Text: + """A model text answer (optionally with a thought summary part and a trailing signature).""" + text: str + thought: str | None = None + signed: bool = True + prompt_tokens: int = 1200 + + +@dataclass +class Call: + """One ``functionCall`` part. The fake mints the id and (for the first call) the signature.""" + name: str + args: dict[str, Any] + + +@dataclass +class Calls: + """A model turn made of ``functionCall`` parts (Gemini 3: signature on the FIRST call only).""" + calls: list[Call] + thought: str | None = "Planning the tool call." + prompt_tokens: int = 1200 + + +@dataclass +class Blocked: + """A candidate stopped by ``finishReason`` (SAFETY / RECITATION / ...) with no content, or, with + ``prompt=True``, a prompt blocked via ``promptFeedback.blockReason`` and no candidates.""" + reason: str = "SAFETY" + prompt: bool = False + + +@dataclass +class GoogleError: + """A Google JSON error body (``google.rpc.Status``) with optional RetryInfo / Retry-After.""" + code: int + status: str + message: str + retry_delay_s: float | None = None + + +@dataclass +class Drop: + """Streaming: emit ``partial`` as one SSE chunk, then close the TLS connection mid-response.""" + partial: str = "Partial answer that never fin" + + +Reply = Text | Calls | Blocked | GoogleError | Drop +Responder = Callable[["Recorded"], "Reply | None"] + + +@dataclass +class Recorded: + method: str + path: str + query: dict[str, list[str]] + headers: dict[str, str] + body: dict[str, Any] | None + version: str = "" + model: str = "" + rpc: str = "" + status: int = 0 + rejection: str | None = None + reply: str = "" + + @property + def stream(self) -> bool: + return self.rpc == "streamGenerateContent" + + @property + def contents(self) -> list[dict[str, Any]]: + return list((self.body or {}).get("contents") or []) + + def parts(self, kind: str) -> list[dict[str, Any]]: + """Every part carrying ``kind`` (e.g. ``functionCall``), in wire order.""" + return [p for c in self.contents for p in c.get("parts") or [] if kind in p] + + def declarations(self) -> dict[str, dict[str, Any]]: + out: dict[str, dict[str, Any]] = {} + for tool in (self.body or {}).get("tools") or []: + for decl in tool.get("functionDeclarations") or []: + out[decl.get("name", "")] = decl + return out + + def all_text(self) -> str: + texts = [p.get("text") or "" for c in self.contents for p in c.get("parts") or []] + system = ((self.body or {}).get("systemInstruction") or {}).get("parts") or [] + return "\n".join(texts + [p.get("text") or "" for p in system]) + + def last_user_text(self) -> str: + for content in reversed(self.contents): + texts = [p["text"] for p in content.get("parts") or [] if isinstance(p.get("text"), str)] + if content.get("role") == "user" and texts: + return "\n".join(texts) + return "" + + +# ── request validation (published contract) ──────────────────────────────────────────────────── +def _validate_schema(node: Any, where: str) -> None: + """``FunctionDeclaration.parameters`` is a proto ``Schema``: unknown keys, list-valued ``type`` + and non-string ``enum`` entries do not parse; ``required`` must name defined properties.""" + schema = _check_fields(node, _SCHEMA_FIELDS, where) + type_ = schema.get("type") + if type_ is not None and (not isinstance(type_, str) or type_.upper() not in _SCHEMA_TYPES): + raise InvalidArgument(f"Invalid value at '{where}.type' (type.googleapis.com/" + f"google.ai.generativelanguage.v1beta.Type), {json.dumps(type_)}") + for i, value in enumerate(schema.get("enum") or []): + if not isinstance(value, str): + raise InvalidArgument(f"Invalid value at '{where}.enum[{i}]' (TYPE_STRING), {json.dumps(value)}") + props = schema.get("properties") or {} + for name, sub in props.items(): + _validate_schema(sub, f"{where}.properties[{name}].value") + for name in schema.get("required") or []: + if name not in props: + raise InvalidArgument(f"{where}.required[{name}]: property is not defined") + if "items" in schema: + _validate_schema(schema["items"], f"{where}.items") + elif isinstance(type_, str) and type_.upper() == "ARRAY": + raise InvalidArgument(f"{where}.items: missing field.") + for i, sub in enumerate(schema.get("anyOf") or []): + _validate_schema(sub, f"{where}.any_of[{i}]") + + +def _validate_json_schema_root(schema: Any, where: str) -> None: + """``parametersJsonSchema`` must describe an object whose properties are the parameters.""" + if not isinstance(schema, dict) or schema.get("type") != "object": + raise InvalidArgument(f"{where}: parameters_json_schema must describe an object (type: object)") + + +def _validate_tools(tools: Any, version: str) -> None: + decl_fields = _DECL_FIELDS_V1BETA if version == "v1beta" else _DECL_FIELDS_V1 + for ti, tool in enumerate(tools if isinstance(tools, list) else []): + tool = _check_fields(tool, _TOOL_FIELDS, f"tools[{ti}]") + for di, decl in enumerate(tool.get("functionDeclarations") or []): + where = f"tools[{ti}].function_declarations[{di}]" + decl = _check_fields(decl, decl_fields, where) + if not _FUNCTION_NAME_RE.match(str(decl.get("name") or "")): + raise InvalidArgument(f"{where}.name: Invalid function name. Must start with a letter or an " + "underscore. Must be alphameric (a-z, A-Z, 0-9), underscores (_), dots (.), " + "colons (:), or dashes (-), with a maximum length of 128.") + if "parameters" in decl and "parametersJsonSchema" in decl: + raise InvalidArgument(f"{where}: parameters and parameters_json_schema are mutually exclusive") + if "parameters" in decl: + _validate_schema(decl["parameters"], f"{where}.parameters") + if "parametersJsonSchema" in decl: + _validate_json_schema_root(decl["parametersJsonSchema"], f"{where}.parameters_json_schema") + + +def _validate_part(part: Any, where: str) -> dict[str, Any]: + part = _check_fields(part, _PART_FIELDS, where) + data = [k for k in part if k in _PART_DATA_FIELDS] + if len(data) != 1: + raise InvalidArgument(f"{where}: a Part must set exactly one data field (oneof 'data'), got {data}") + if "functionCall" in part: + fc = _check_fields(part["functionCall"], _FUNCTION_CALL_FIELDS, f"{where}.function_call") + if not fc.get("name") or not isinstance(fc.get("args", {}), dict): + raise InvalidArgument(f"{where}.function_call: name is required and args must be an object") + if "functionResponse" in part: + fr = _check_fields(part["functionResponse"], _FUNCTION_RESPONSE_FIELDS, f"{where}.function_response") + if not fr.get("name") or not isinstance(fr.get("response"), dict): + raise InvalidArgument(f"{where}.function_response: name is required and response must be an object") + sig = part.get("thoughtSignature") + if sig is not None and not (isinstance(sig, str) and sig): + raise InvalidArgument(f"{where}.thought_signature: invalid bytes value") + return part + + +def _check_pairing(contents: list[dict[str, Any]], gemini3: bool) -> None: + """Every functionCall turn is answered by the NEXT content with one functionResponse per call + (same name, same id on Gemini 3); a functionResponse turn must follow a functionCall turn.""" + for i, content in enumerate(contents): + calls = [p["functionCall"] for p in content["parts"] if "functionCall" in p] + responses = [p["functionResponse"] for p in content["parts"] if "functionResponse" in p] + if responses: + prev_calls = [p["functionCall"] for p in contents[i - 1]["parts"] if "functionCall" in p] if i else [] + if not prev_calls: + raise InvalidArgument("Please ensure that function response turn comes immediately after a " + "function call turn.") + if not calls: + continue + nxt = contents[i + 1]["parts"] if i + 1 < len(contents) else [] + answered = [p["functionResponse"] for p in nxt if "functionResponse" in p] + if i + 1 < len(contents) and len(answered) != len(calls): + raise InvalidArgument("Please ensure that the number of function response parts is equal to the " + "number of function call parts of the function call turn.") + for call, resp in zip(calls, answered): + if call.get("name") != resp.get("name"): + raise InvalidArgument(f"functionResponse name {resp.get('name')!r} does not match functionCall " + f"name {call.get('name')!r} in contents[{i + 1}]") + if gemini3 and answered: + if sorted(str(c.get("id")) for c in calls) != sorted(str(r.get("id")) for r in answered): + raise InvalidArgument(f"functionResponse ids in contents[{i + 1}] do not match the functionCall " + "ids of the function call turn.") + + +def _current_turn_start(contents: list[dict[str, Any]]) -> int: + """Index of the newest user content with standard (non-functionResponse) content.""" + for i in range(len(contents) - 1, -1, -1): + c = contents[i] + if c["role"] == "user" and any("functionResponse" not in p for p in c["parts"]): + return i + return 0 + + +def _check_signatures(contents: list[dict[str, Any]], issued: set[str]) -> None: + """Gemini 3: the FIRST functionCall part of each step of the current turn must carry a + thoughtSignature Google issued (or a documented dummy); a forged one is corrupted.""" + for i in range(_current_turn_start(contents), len(contents)): + c = contents[i] + fc_parts = [p for p in c["parts"] if "functionCall" in p] if c["role"] == "model" else [] + if not fc_parts: + continue + sig = fc_parts[0].get("thoughtSignature") + if not sig: + raise InvalidArgument(f"Function call `{fc_parts[0]['functionCall']['name']}` in the `{i}.` content " + "block is missing a `thought_signature`.") + if sig not in issued and sig not in SKIP_SIGNATURES: + raise InvalidArgument("Corrupted thought signature.") + + +def validate_generate_request(body: Any, version: str, model: str, issued: set[str]) -> None: + body = _check_fields(body, _REQUEST_FIELDS, "") + contents = body.get("contents") + if not isinstance(contents, list) or not contents: + raise InvalidArgument("* GenerateContentRequest.contents: contents is not specified") + norm: list[dict[str, Any]] = [] + for ci, content in enumerate(contents): + content = _check_fields(content, _CONTENT_FIELDS, f"contents[{ci}]") + role = content.get("role", "user") + if role not in _ROLES: + raise InvalidArgument(f"Please use a valid role: user, model. (contents[{ci}].role={role!r})") + parts = content.get("parts") + if not isinstance(parts, list) or not parts: + raise InvalidArgument(f"* GenerateContentRequest.contents[{ci}].parts: contents.parts must not be empty.") + norm.append({"role": role, "parts": [_validate_part(p, f"contents[{ci}].parts[{pi}]") + for pi, p in enumerate(parts)]}) + for a, b in zip(norm, norm[1:]): + if a["role"] == b["role"]: + raise InvalidArgument("Please ensure that multiturn requests alternate between user and model.") + if norm[-1]["role"] != "user": + raise InvalidArgument("Please ensure that single turn requests end with a user role or the role field " + "is empty.") + gemini3 = bool(re.match(r"gemini-([3-9]|\d\d)", model)) + _check_pairing(norm, gemini3) + if gemini3: + _check_signatures(norm, issued) + if "systemInstruction" in body: + system = _check_fields(body["systemInstruction"], _CONTENT_FIELDS, "system_instruction") + for pi, p in enumerate(system.get("parts") or []): + _validate_part(p, f"system_instruction.parts[{pi}]") + _validate_tools(body.get("tools"), version) + generation = _check_fields(body.get("generationConfig") or {}, _GENERATION_FIELDS, "generation_config") + if "thinkingConfig" in generation: + _check_fields(generation["thinkingConfig"], _THINKING_FIELDS, "generation_config.thinking_config") + + +# ── response building (GenerateContentResponse) ──────────────────────────────────────────────── +def _usage(prompt_tokens: int, output_tokens: int = 24, thought_tokens: int = 16) -> dict[str, Any]: + return {"promptTokenCount": prompt_tokens, "candidatesTokenCount": output_tokens, + "thoughtsTokenCount": thought_tokens, "totalTokenCount": prompt_tokens + output_tokens + thought_tokens} + + +def _response(parts: list[dict[str, Any]] | None, finish: str | None, usage: dict[str, Any] | None, + extra: dict[str, Any] | None = None) -> dict[str, Any]: + cand: dict[str, Any] = {"index": 0} + if parts is not None: + cand["content"] = {"role": "model", "parts": parts} + if finish: + cand["finishReason"] = finish + out: dict[str, Any] = {"candidates": [cand], "modelVersion": MODEL_ID, + "responseId": secrets.token_urlsafe(12)} + if usage: + out["usageMetadata"] = usage + out.update(extra or {}) + return out + + +def _chunks(text: str, n: int = 3) -> list[str]: + step = max(1, -(-len(text) // n)) + return [text[i:i + step] for i in range(0, len(text), step)] or [""] + + +# ── TLS material ──────────────────────────────────────────────────────────────────────────────── +def _write_tls_material(directory: Path) -> tuple[Path, Path, Path]: + """Throwaway CA + a leaf for the Google host. Returns (ca_pem, leaf_cert_pem, leaf_key_pem).""" + from cryptography import x509 + from cryptography.hazmat.primitives import hashes, serialization + from cryptography.hazmat.primitives.asymmetric import ec + from cryptography.x509.oid import ExtendedKeyUsageOID, NameOID + + now = datetime.datetime.now(datetime.timezone.utc) + ca_key = ec.generate_private_key(ec.SECP256R1()) + ca_name = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, "hermes-e2e gemini fake CA")]) + ca = (x509.CertificateBuilder().subject_name(ca_name).issuer_name(ca_name) + .public_key(ca_key.public_key()).serial_number(x509.random_serial_number()) + .not_valid_before(now - datetime.timedelta(minutes=5)).not_valid_after(now + datetime.timedelta(days=1)) + .add_extension(x509.BasicConstraints(ca=True, path_length=0), critical=True) + .add_extension(x509.KeyUsage(digital_signature=True, key_cert_sign=True, crl_sign=True, + content_commitment=False, key_encipherment=False, data_encipherment=False, + key_agreement=False, encipher_only=False, decipher_only=False), critical=True) + .add_extension(x509.SubjectKeyIdentifier.from_public_key(ca_key.public_key()), critical=False) + .sign(ca_key, hashes.SHA256())) + leaf_key = ec.generate_private_key(ec.SECP256R1()) + leaf = (x509.CertificateBuilder() + .subject_name(x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, GEMINI_HOST)])) + .issuer_name(ca_name).public_key(leaf_key.public_key()).serial_number(x509.random_serial_number()) + .not_valid_before(now - datetime.timedelta(minutes=5)).not_valid_after(now + datetime.timedelta(days=1)) + .add_extension(x509.SubjectAlternativeName([x509.DNSName(GEMINI_HOST), + x509.IPAddress(ipaddress.ip_address("127.0.0.1"))]), + critical=False) + .add_extension(x509.BasicConstraints(ca=False, path_length=None), critical=True) + .add_extension(x509.ExtendedKeyUsage([ExtendedKeyUsageOID.SERVER_AUTH]), critical=False) + .add_extension(x509.AuthorityKeyIdentifier.from_issuer_public_key(ca_key.public_key()), critical=False) + .sign(ca_key, hashes.SHA256())) + directory.mkdir(parents=True, exist_ok=True) + ca_pem, cert_pem, key_pem = directory / "ca.pem", directory / "leaf.pem", directory / "leaf.key" + ca_pem.write_bytes(ca.public_bytes(serialization.Encoding.PEM)) + cert_pem.write_bytes(leaf.public_bytes(serialization.Encoding.PEM)) + key_pem.write_bytes(leaf_key.private_bytes(serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, + serialization.NoEncryption())) + return ca_pem, cert_pem, key_pem + + +# ── the fake ──────────────────────────────────────────────────────────────────────────────────── +class GeminiFake: + """Loopback CONNECT proxy + TLS-terminated fake Gemini API. + + ``script`` replies are consumed in order by every valid generate call that ``route`` (optional + responder, e.g. for compaction summaries) does not claim. An exhausted script answers with a + Google 500 so an unexpected extra call is visible instead of hanging. + """ + + def __init__(self, workdir: Path, script: list[Reply] | None = None, *, route: Responder | None = None, + api_key: str = API_KEY) -> None: + self.workdir = workdir + self.api_key = api_key + self.script: list[Reply] = list(script or []) + self.route = route + self.requests: list[Recorded] = [] + self.refused_hosts: list[str] = [] + self.issued_signatures: list[str] = [] + self.call_signatures: dict[str, str] = {} # functionCall id -> signature minted on its part + self._lock = threading.Lock() + self._counter = 0 + self.ca_pem, cert, key = _write_tls_material(workdir / "tls") + self._tls = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + self._tls.load_cert_chain(cert, key) + self._server = ThreadingHTTPServer(("127.0.0.1", 0), self._handler_class()) + self._server.daemon_threads = True + self._thread = threading.Thread(target=self._server.serve_forever, name="gemini-fake", daemon=True) + + # lifecycle --------------------------------------------------------------------------------- + def __enter__(self) -> "GeminiFake": + self._thread.start() + return self + + def __exit__(self, *exc: Any) -> None: + self._server.shutdown() + self._server.server_close() + self._thread.join(timeout=10) + + @property + def proxy_url(self) -> str: + return f"http://127.0.0.1:{self._server.server_address[1]}" + + def child_env(self) -> dict[str, str]: + """Env for the ``hermes`` child: route HTTPS through the proxy and trust the fake CA.""" + ca = str(self.ca_pem) + return {"HTTPS_PROXY": self.proxy_url, "https_proxy": self.proxy_url, + "HTTP_PROXY": self.proxy_url, "http_proxy": self.proxy_url, + "NO_PROXY": "", "no_proxy": "", + "SSL_CERT_FILE": ca, "REQUESTS_CA_BUNDLE": ca, "CURL_CA_BUNDLE": ca} + + # views ------------------------------------------------------------------------------------ + def generate_calls(self) -> list[Recorded]: + with self._lock: + return [r for r in self.requests if r.rpc in {"generateContent", "streamGenerateContent"}] + + def rejections(self) -> list[str]: + with self._lock: + return [f"{r.rpc} -> {r.status}: {r.rejection}" for r in self.requests if r.rejection] + + def main_calls(self) -> list[Recorded]: + """Generate calls answered from ``script`` (not claimed by ``route``).""" + return [r for r in self.generate_calls() if r.reply.startswith("script:")] + + # dispatch --------------------------------------------------------------------------------- + def _mint_signature(self) -> str: + sig = base64.b64encode(b"\x12\x34gemini-e2e-sig:" + secrets.token_bytes(24)).decode() + with self._lock: + self.issued_signatures.append(sig) + return sig + + def _mint_call_id(self) -> str: + with self._lock: + self._counter += 1 + return f"fc-{self._counter:04d}-{secrets.token_hex(3)}" + + def _next_reply(self, rec: Recorded) -> Reply: + if self.route is not None and (routed := self.route(rec)) is not None: + rec.reply = f"route:{type(routed).__name__}" + return routed + with self._lock: + reply = self.script.pop(0) if self.script else GoogleError(500, "INTERNAL", "fake script exhausted") + rec.reply = f"script:{type(reply).__name__}" + return reply + + def _build(self, reply: Reply) -> list[dict[str, Any]]: + """Stream events for a reply (a unary response is their merge).""" + builders: dict[type, Callable[[Any], list[dict[str, Any]]]] = { + Text: self._text_events, Calls: self._call_events, Blocked: self._blocked_events, + } + return builders[type(reply)](reply) + + def _text_events(self, reply: Text) -> list[dict[str, Any]]: + events: list[dict[str, Any]] = [] + if reply.thought: + events.append(_response([{"text": reply.thought, "thought": True}], None, None)) + pieces = _chunks(reply.text) + for i, piece in enumerate(pieces): + part: dict[str, Any] = {"text": piece} + last = i == len(pieces) - 1 + if last and reply.signed: + part["thoughtSignature"] = self._mint_signature() + events.append(_response([part], "STOP" if last else None, _usage(reply.prompt_tokens) if last else None)) + return events + + def _call_events(self, reply: Calls) -> list[dict[str, Any]]: + events: list[dict[str, Any]] = [] + if reply.thought: + events.append(_response([{"text": reply.thought, "thought": True}], None, None)) + parts = [] + for i, call in enumerate(reply.calls): + call_id = self._mint_call_id() + part: dict[str, Any] = {"functionCall": {"id": call_id, "name": call.name, "args": call.args}} + if i == 0: + part["thoughtSignature"] = self._mint_signature() + with self._lock: + self.call_signatures[call_id] = part["thoughtSignature"] + parts.append(part) + events.append(_response(parts, "STOP", _usage(reply.prompt_tokens))) + return events + + @staticmethod + def _blocked_events(reply: Blocked) -> list[dict[str, Any]]: + if reply.prompt: + return [{"promptFeedback": {"blockReason": reply.reason, "safetyRatings": [ + {"category": "HARM_CATEGORY_DANGEROUS_CONTENT", "probability": "HIGH", "blocked": True}]}, + "usageMetadata": {"promptTokenCount": 900, "totalTokenCount": 900}, "modelVersion": MODEL_ID}] + ratings = [{"category": "HARM_CATEGORY_DANGEROUS_CONTENT", "probability": "HIGH", "blocked": True}] + extra = {"citationMetadata": {"citationSources": [{"startIndex": 0, "endIndex": 40, + "uri": "https://example.com/source"}]}} + cand_extra = ratings if reply.reason == "SAFETY" else None + resp = _response(None, reply.reason, _usage(900, 0, 0)) + if cand_extra: + resp["candidates"][0]["safetyRatings"] = cand_extra + if reply.reason == "RECITATION": + resp["candidates"][0].update(extra) + return [resp] + + @staticmethod + def merge_events(events: list[dict[str, Any]]) -> dict[str, Any]: + """Unary ``generateContent`` body = the stream's chunks folded into one candidate.""" + if not events or "candidates" not in events[-1]: + return events[-1] if events else {} + parts: list[dict[str, Any]] = [] + for ev in events: + parts.extend(((ev.get("candidates") or [{}])[0].get("content") or {}).get("parts") or []) + final = json.loads(json.dumps(events[-1])) + if parts: + final["candidates"][0]["content"] = {"role": "model", "parts": parts} + return final + + # HTTP handler ----------------------------------------------------------------------------- + def _handler_class(self) -> type[BaseHTTPRequestHandler]: + fake = self + + class Handler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + tunneled = False + + def log_message(self, format: str, *args: Any) -> None: # noqa: A002 - base signature + return + + def do_CONNECT(self) -> None: # noqa: N802 - http.server naming + host = self.path.split(":", 1)[0].lower() + if self.tunneled or host != GEMINI_HOST: + with fake._lock: + fake.refused_hosts.append(host) + self.send_response(403, "Forbidden by hermes e2e fake proxy") + self.send_header("Content-Length", "0") + self.end_headers() + self.close_connection = True + return + self.send_response(200, "Connection established") + self.end_headers() + self.wfile.flush() + try: + tls = fake._tls.wrap_socket(self.connection, server_side=True) + except (ssl.SSLError, OSError): + self.close_connection = True + return + self.connection = self.request = tls + self.rfile = tls.makefile("rb") + self.wfile = tls.makefile("wb") + self.tunneled = True + self.close_connection = False + + def _refuse_plain(self) -> None: + with fake._lock: + fake.refused_hosts.append(urlsplit(self.path).hostname or self.headers.get("Host", "?")) + self.send_response(403) + self.send_header("Content-Length", "0") + self.end_headers() + + def do_GET(self) -> None: # noqa: N802 + if not self.tunneled: + return self._refuse_plain() + fake._handle(self, "GET") + + def do_POST(self) -> None: # noqa: N802 + if not self.tunneled: + return self._refuse_plain() + fake._handle(self, "POST") + + return Handler + + def _send_json(self, h: BaseHTTPRequestHandler, rec: Recorded, status: int, body: dict[str, Any], + headers: dict[str, str] | None = None) -> None: + rec.status = status + data = json.dumps(body).encode() + h.send_response(status) + h.send_header("Content-Type", "application/json; charset=UTF-8") + for k, v in (headers or {}).items(): + h.send_header(k, v) + h.send_header("Content-Length", str(len(data))) + h.end_headers() + h.wfile.write(data) + h.wfile.flush() + + def _error(self, h: BaseHTTPRequestHandler, rec: Recorded, err: GoogleError) -> None: + body: dict[str, Any] = {"error": {"code": err.code, "message": err.message, "status": err.status}} + headers = {} + if err.retry_delay_s is not None: + body["error"]["details"] = [{"@type": "type.googleapis.com/google.rpc.RetryInfo", + "retryDelay": f"{err.retry_delay_s:g}s"}] + headers["Retry-After"] = f"{err.retry_delay_s:g}" + if err.code == 400: + rec.rejection = rec.rejection or err.message + self._send_json(h, rec, err.code, body, headers) + + def _handle(self, h: BaseHTTPRequestHandler, method: str) -> None: + url = urlsplit(h.path) + length = int(h.headers.get("Content-Length") or 0) + raw = h.rfile.read(length) if length else b"" + try: + body = json.loads(raw) if raw else None + except ValueError: + body = None + rec = Recorded(method, url.path, parse_qs(url.query), {k.lower(): v for k, v in h.headers.items()}, body) + with self._lock: + self.requests.append(rec) + match = _PATH_RE.match(url.path) + if method == "GET" or not match: + return self._error(h, rec, GoogleError(404, "NOT_FOUND", f"fake has no route for {method} {url.path}")) + rec.version, rec.model, rec.rpc = match["version"], match["model"], match["method"] + key = rec.headers.get("x-goog-api-key") or (rec.query.get("key") or [""])[0] + if key != self.api_key: + rec.rejection = "API key not valid" + return self._error(h, rec, GoogleError(400, "INVALID_ARGUMENT", + "API key not valid. Please pass a valid API key.")) + if rec.stream and (rec.query.get("alt") or [""])[0] != "sse": + rec.rejection = "streamGenerateContent without alt=sse" + try: + with self._lock: + issued = set(self.issued_signatures) + validate_generate_request(body, rec.version, rec.model, issued) + except InvalidArgument as exc: + rec.rejection = str(exc) + return self._error(h, rec, GoogleError(400, "INVALID_ARGUMENT", str(exc))) + reply = self._next_reply(rec) + if isinstance(reply, GoogleError): + return self._error(h, rec, reply) + if isinstance(reply, Drop): + return self._drop(h, rec, reply) + events = self._build(reply) + if not rec.stream: + return self._send_json(h, rec, 200, self.merge_events(events)) + self._stream(h, rec, events) + + def _stream(self, h: BaseHTTPRequestHandler, rec: Recorded, events: list[dict[str, Any]]) -> None: + rec.status = 200 + h.send_response(200) + h.send_header("Content-Type", "text/event-stream") + h.send_header("Transfer-Encoding", "chunked") + h.end_headers() + for ev in events: + self._write_chunk(h, f"data: {json.dumps(ev)}\r\n\r\n".encode()) + h.wfile.write(b"0\r\n\r\n") + h.wfile.flush() + + @staticmethod + def _write_chunk(h: BaseHTTPRequestHandler, data: bytes) -> None: + h.wfile.write(f"{len(data):x}\r\n".encode() + data + b"\r\n") + h.wfile.flush() + + def _drop(self, h: BaseHTTPRequestHandler, rec: Recorded, reply: Drop) -> None: + rec.status = 200 + h.send_response(200) + h.send_header("Content-Type", "text/event-stream") + h.send_header("Transfer-Encoding", "chunked") + h.end_headers() + self._write_chunk(h, f"data: {json.dumps(_response([{'text': reply.partial}], None, None))}\r\n\r\n".encode()) + time.sleep(0.05) # let the partial chunk reach the client before the reset + h.close_connection = True + try: + h.connection.setsockopt(socket.SOL_SOCKET, socket.SO_LINGER, struct.pack("ii", 1, 0)) + h.connection.close() + except OSError: + pass diff --git a/tests/fakes/providers/vertex.py b/tests/fakes/providers/vertex.py new file mode 100644 index 000000000000..6936388347b0 --- /dev/null +++ b/tests/fakes/providers/vertex.py @@ -0,0 +1,709 @@ +"""Loopback fake of Google Vertex AI (Gemini behind the OpenAI-compatible endpoint) + Google OAuth. + +Two external boundaries Hermes does not own, both faked for real: + +* **Google OAuth2 token endpoint** (``POST /token`` on plain loopback HTTP). A generated + service-account JSON names it as ``token_uri``, so the REAL ``google-auth`` library signs an + RS256 JWT assertion with the SA private key and exchanges it. The fake verifies the signature + against the SA public key and the claims Google checks (``iss``/``aud``/``scope``/``iat``/``exp``) + and mints ``ya29.``-style access tokens with a scripted ``expires_in``. +* **Vertex AI** at ``https://{region}-aiplatform.googleapis.com/v1beta1/projects/{project}/locations/ + {region}/endpoints/openapi/chat/completions``. Hermes has no base-URL override for Vertex, so the + fake is an HTTPS ``CONNECT`` proxy that terminates TLS with a leaf cert signed by a generated CA: + the child trusts it through the standard ``SSL_CERT_FILE`` and reaches it through the standard + ``HTTPS_PROXY`` (the corporate-proxy channel Hermes documents). CONNECTs to any other host are + refused (recorded), so nothing can leak to the real network. + +Every Vertex request is recorded (host, path, headers, body) and validated against the published +OpenAI-compatibility contract before a scripted response is chosen; a malformed request gets the +error Vertex returns (``[{"error": {"code", "message", "status"}}]``) and is marked ``rejected``. +Checks: URL scheme (project/location path segments), a live minted bearer, the ``google/`` +publisher form, tool-call/tool-result pairing, JSON-object function arguments, function-name rules, +and Gemini 3 thought signatures (``extra_content.google.thought_signature``) on the first function +call of every step in the current turn, byte-identical to one the fake issued. +""" + +from __future__ import annotations + +import base64 +import datetime as _dt +import json +import re +import secrets +import socketserver +import ssl +import threading +import time +import urllib.parse +from dataclasses import dataclass +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from pathlib import Path +from typing import Any, Callable, Union + +from cryptography import x509 +from cryptography.hazmat.primitives import hashes, serialization +from cryptography.hazmat.primitives.asymmetric import ec, padding, rsa +from cryptography.x509.oid import ExtendedKeyUsageOID, NameOID + +CLOUD_PLATFORM_SCOPE = "https://www.googleapis.com/auth/cloud-platform" +JWT_BEARER_GRANT = "urn:ietf:params:oauth:grant-type:jwt-bearer" +# Google's token endpoint requires this audience whatever host the SA file's token_uri names. +GOOGLE_TOKEN_AUDIENCE = "https://oauth2.googleapis.com/token" +UNAUTHENTICATED_MESSAGE = ( + "Request had invalid authentication credentials. Expected OAuth 2 access token, login cookie or " + "other valid authentication credential. See https://developers.google.com/identity/sign-in/web/devconsole-project.") +MISSING_SIGNATURE_MESSAGE = ( + "Function call is missing a thought_signature in functionCall parts. This is required for tools to work " + "correctly, and missing thought_signature may lead to degraded model performance. Additional data, function " + "call `default_api:{name}` , position {pos}. Please refer to https://ai.google.dev/gemini-api/docs/" + "thought-signatures for more details.") +# HTTP status -> google.rpc canonical status name (https://cloud.google.com/apis/design/errors). +GRPC_STATUS = {400: "INVALID_ARGUMENT", 401: "UNAUTHENTICATED", 403: "PERMISSION_DENIED", 404: "NOT_FOUND", + 429: "RESOURCE_EXHAUSTED", 500: "INTERNAL", 503: "UNAVAILABLE", 504: "DEADLINE_EXCEEDED"} +_MODEL_RE = re.compile(r"^google/gemini-[0-9a-z.\-]+$") +# FunctionDeclaration.name: letter/underscore first, then [a-zA-Z0-9_.:-], at most 64 chars. +_FUNC_NAME_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_.:\-]{0,63}$") +_ROLES = frozenset({"system", "developer", "user", "assistant", "tool"}) + + +# Scripted responses -------------------------------------------------------------------------------- + + +@dataclass +class Say: + """A final assistant answer streamed in chunks.""" + + text: str + prompt_tokens: int | None = None + completion_tokens: int = 20 + chunk_chars: int = 16 + expire_tokens_after: bool = False # the ~1h boundary passes right after this response + + +@dataclass +class Call: + """One assistant step issuing function calls; Gemini 3 signs the first call of each step.""" + + calls: list[tuple[str, dict[str, Any]]] + text: str | None = None + signed: bool = True + prompt_tokens: int | None = None + expire_tokens_after: bool = False + + +@dataclass +class Fail: + """A Vertex error response; ``list_body`` is the openapi endpoint's list-wrapped envelope.""" + + status: int + message: str + list_body: bool = True + retry_after: float | None = None + + +@dataclass +class Drop: + """Open the SSE stream, send ``text[:after_chars]``, then reset the TLS connection mid-body.""" + + text: str + after_chars: int = 10 + + +Response = Union[Say, Call, Fail, Drop] +Responder = Callable[[dict[str, Any]], Response] + + +# Certificates and service account -------------------------------------------------------------------- + + +def _name(cn: str) -> x509.Name: + return x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, cn)]) + + +def _write_pem(path: Path, data: bytes) -> Path: + path.write_bytes(data) + return path + + +def make_tls_material(root: Path, hosts: list[str]) -> tuple[Path, Path, Path]: + """A throwaway CA (PEM for ``SSL_CERT_FILE``) and a leaf for ``hosts`` signed by it.""" + root.mkdir(parents=True, exist_ok=True) + now = _dt.datetime.now(_dt.timezone.utc) + ca_key = ec.generate_private_key(ec.SECP256R1()) + ca_ski = x509.SubjectKeyIdentifier.from_public_key(ca_key.public_key()) + ca_cert = ( + x509.CertificateBuilder().subject_name(_name("hermes-e2e fake Google CA")).issuer_name(_name("hermes-e2e fake Google CA")) + .public_key(ca_key.public_key()).serial_number(x509.random_serial_number()) + .not_valid_before(now - _dt.timedelta(minutes=5)).not_valid_after(now + _dt.timedelta(days=1)) + .add_extension(x509.BasicConstraints(ca=True, path_length=0), critical=True) + .add_extension(x509.KeyUsage(digital_signature=True, key_cert_sign=True, crl_sign=True, content_commitment=False, + key_encipherment=False, data_encipherment=False, key_agreement=False, + encipher_only=False, decipher_only=False), critical=True) + .add_extension(ca_ski, critical=False) + .sign(ca_key, hashes.SHA256())) + leaf_key = ec.generate_private_key(ec.SECP256R1()) + leaf_cert = ( + x509.CertificateBuilder().subject_name(_name(hosts[0])).issuer_name(ca_cert.subject) + .public_key(leaf_key.public_key()).serial_number(x509.random_serial_number()) + .not_valid_before(now - _dt.timedelta(minutes=5)).not_valid_after(now + _dt.timedelta(days=1)) + .add_extension(x509.SubjectAlternativeName([x509.DNSName(h) for h in hosts]), critical=False) + .add_extension(x509.BasicConstraints(ca=False, path_length=None), critical=True) + .add_extension(x509.ExtendedKeyUsage([ExtendedKeyUsageOID.SERVER_AUTH]), critical=False) + .add_extension(x509.AuthorityKeyIdentifier.from_issuer_subject_key_identifier(ca_ski), critical=False) + .sign(ca_key, hashes.SHA256())) + ca_pem = _write_pem(root / "fake-google-ca.pem", ca_cert.public_bytes(serialization.Encoding.PEM)) + cert_pem = _write_pem(root / "leaf.pem", leaf_cert.public_bytes(serialization.Encoding.PEM)) + key_pem = _write_pem(root / "leaf.key", leaf_key.private_bytes( + serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption())) + return ca_pem, cert_pem, key_pem + + +@dataclass +class ServiceAccount: + path: Path + client_email: str + project_id: str + private_key_id: str + public_key: rsa.RSAPublicKey + + +def make_service_account(path: Path, token_uri: str, project_id: str) -> ServiceAccount: + """A service-account key file shaped like the one the Cloud console downloads.""" + key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + email = f"hermes-e2e@{project_id}.iam.gserviceaccount.com" + kid = secrets.token_hex(20) + info = { + "type": "service_account", "project_id": project_id, "private_key_id": kid, + "private_key": key.private_bytes(serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, + serialization.NoEncryption()).decode(), + "client_email": email, "client_id": str(secrets.randbelow(10**20)), + "auth_uri": "https://accounts.google.com/o/oauth2/auth", "token_uri": token_uri, + "auth_provider_x509_cert_url": "https://www.googleapis.com/oauth2/v1/certs", + "universe_domain": "googleapis.com", + } + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(json.dumps(info, indent=2), encoding="utf-8") + return ServiceAccount(path, email, project_id, kid, key.public_key()) + + +def _b64url_decode(part: str) -> bytes: + return base64.urlsafe_b64decode(part + "=" * (-len(part) % 4)) + + +def verify_jwt_assertion(assertion: str, sa: ServiceAccount) -> tuple[dict[str, Any] | None, str]: + """RS256 signature + the claims Google's token endpoint enforces; (claims, "") or (None, reason).""" + try: + head_b64, body_b64, sig_b64 = assertion.split(".") + header, claims = json.loads(_b64url_decode(head_b64)), json.loads(_b64url_decode(body_b64)) + sa.public_key.verify(_b64url_decode(sig_b64), f"{head_b64}.{body_b64}".encode(), padding.PKCS1v15(), hashes.SHA256()) + except Exception as exc: # noqa: BLE001 - any parse/verify failure is Google's "Invalid JWT Signature." + return None, f"Invalid JWT Signature. ({type(exc).__name__})" + now = time.time() + problems = { + "alg": header.get("alg") != "RS256", + "kid": header.get("kid") not in (None, sa.private_key_id), + "iss": claims.get("iss") != sa.client_email, + "aud": claims.get("aud") != GOOGLE_TOKEN_AUDIENCE, + "scope": CLOUD_PLATFORM_SCOPE not in str(claims.get("scope", "")).split(), + "iat": not isinstance(claims.get("iat"), int) or abs(claims["iat"] - now) > 300, + "exp": not isinstance(claims.get("exp"), int) or not 0 < claims["exp"] - claims.get("iat", 0) <= 3600, + } + bad = sorted(k for k, v in problems.items() if v) + return (None, f"Invalid JWT: bad claims {bad}") if bad else (claims, "") + + +# Request validation ---------------------------------------------------------------------------------- + + +def _first_signature(tool_call: dict[str, Any]) -> Any: + extra = tool_call.get("extra_content") + google = extra.get("google") if isinstance(extra, dict) else None + return google.get("thought_signature") if isinstance(google, dict) else None + + +def _check_tool_call_shape(tc: Any) -> str | None: + if not isinstance(tc, dict) or tc.get("type", "function") != "function" or not isinstance(tc.get("function"), dict): + return "Invalid value at 'messages[].tool_calls[]': expected a function tool call" + fn = tc["function"] + if not isinstance(tc.get("id"), str) or not tc["id"]: + return "tool_calls[].id must be a non-empty string" + try: + args = json.loads(fn.get("arguments") or "{}") + except (TypeError, json.JSONDecodeError): + return f"Invalid JSON payload in function call arguments for `{fn.get('name')}`" + return None if isinstance(args, dict) else "function call arguments must be a JSON object (google.protobuf.Struct)" + + +def _check_pairing(messages: list[dict[str, Any]]) -> str | None: + """Every function-call turn is answered by exactly its function responses, immediately after.""" + i = 0 + while i < len(messages): + msg = messages[i] + if msg.get("role") == "tool": + return ("Please ensure that function response turn comes immediately after a function call turn. " + f"(orphan tool message at index {i}, tool_call_id={msg.get('tool_call_id')!r})") + calls = msg.get("tool_calls") if msg.get("role") == "assistant" else None + if not calls: + i += 1 + continue + expected = [tc.get("id") for tc in calls] + j = i + 1 + answered: list[Any] = [] + while j < len(messages) and messages[j].get("role") == "tool": + answered.append(messages[j].get("tool_call_id")) + j += 1 + if sorted(map(str, answered)) != sorted(map(str, expected)): + return ("Please ensure that the number of function response parts is equal to the number of function " + f"call parts of the function call turn. (calls {expected} at index {i}, responses {answered})") + i = j + return None + + +def _check_signatures(messages: list[dict[str, Any]], issued: set[str]) -> str | None: + """Gemini 3: the first call of every step in the current turn carries a signature we issued; + any replayed signature (current or historical) must be byte-identical to one we issued.""" + last_user = max((i for i, m in enumerate(messages) if m.get("role") == "user"), default=-1) + step = 0 + for i, msg in enumerate(messages): + calls = msg.get("tool_calls") if msg.get("role") == "assistant" else None + if not calls: + continue + for pos, tc in enumerate(calls): + sig = _first_signature(tc) + if sig is not None and sig not in issued: + return "Corrupted thought signature." + if i > last_user and pos == 0 and sig is None: + return MISSING_SIGNATURE_MESSAGE.format(name=tc.get("function", {}).get("name"), pos=step + 1) + step += 1 + return None + + +def _merged_anyof_error(node: Any, where: str) -> str | None: + """Google's OpenAI->FunctionDeclaration translator merges ``anyOf`` branches into one node; a + merged node carrying ``items`` under a non-array type is rejected (#109115 evidence).""" + if isinstance(node, dict): + branches = node.get("anyOf") + if isinstance(branches, list) and len(branches) > 1: + types = {b.get("type") for b in branches if isinstance(b, dict)} + if len(types) > 1 and any(isinstance(b, dict) and "items" in b for b in branches): + return (f"functionDeclaration `{where}` schema specified incorrect schema type field. " + "For schema with items, schema type should be ARRAY.") + for key, child in node.items(): + found = _merged_anyof_error(child, f"{where}.{key}") + if found: + return found + elif isinstance(node, list): + for child in node: + found = _merged_anyof_error(child, where) + if found: + return found + return None + + +def _check_tools(tools: Any) -> str | None: + if tools is None: + return None + if not isinstance(tools, list): + return "Invalid value at 'tools': expected a list" + seen: set[str] = set() + for tool in tools: + fn = tool.get("function") if isinstance(tool, dict) and tool.get("type") == "function" else None + if not isinstance(fn, dict): + return "Invalid value at 'tools[]': only function tools are supported" + name = str(fn.get("name", "")) + if not _FUNC_NAME_RE.match(name): + return f"Invalid function name `{name}`: must start with a letter or underscore, [a-zA-Z0-9_.:-], max 64" + if name in seen: + return f"Duplicate function declaration found: {name}" + seen.add(name) + params = fn.get("parameters") + if params is not None and (not isinstance(params, dict) or params.get("type", "object") != "object"): + return f"functionDeclaration `{name}` parameters must be an OBJECT schema" + found = _merged_anyof_error(params, f"{name}.parameters") + if found: + return f"Unable to submit request because `{name}` {found}" + return None + + +def validate_chat_body(body: dict[str, Any], issued_signatures: set[str]) -> str | None: + """The first contract violation Vertex would 400 on, or None.""" + if not _MODEL_RE.match(str(body.get("model", ""))): + return f"Invalid model name {body.get('model')!r}: the OpenAI-compatible endpoint expects 'google/'" + messages = body.get("messages") + if not isinstance(messages, list) or not messages: + return "* GenerateContentRequest.contents: contents is not specified" + for idx, msg in enumerate(messages): + if not isinstance(msg, dict) or msg.get("role") not in _ROLES: + return f"Invalid value at 'messages[{idx}].role'" + if msg.get("role") == "tool" and not msg.get("tool_call_id"): + return f"messages[{idx}]: a tool message requires tool_call_id" + for tc in msg.get("tool_calls") or []: + bad = _check_tool_call_shape(tc) + if bad: + return f"messages[{idx}]: {bad}" + return _check_tools(body.get("tools")) or _check_pairing(messages) or _check_signatures(messages, issued_signatures) + + +# Server ---------------------------------------------------------------------------------------------- + + +@dataclass +class _Token: + value: str + expires_at: float + + +@dataclass +class TokenPolicy: + expires_in: int = 3600 + error: tuple[int, str, str] | None = None # (status, error, error_description) for every exchange + reject_bearers: bool = False # Vertex refuses every bearer (disabled SA / revoked grant) + + +class FakeVertex: + """OAuth token endpoint + TLS-terminating CONNECT proxy serving Vertex. Context manager.""" + + def __init__(self, root: Path, *, project: str, region: str, sa_project: str | None = None, + script: list[Response] | Responder | None = None, aux: Responder | None = None, + default_text: str = "ok", prompt_tokens_fn: Callable[[dict[str, Any]], int] | None = None) -> None: + self.root, self.project, self.region = root, project, region + self.host = "aiplatform.googleapis.com" if region == "global" else f"{region}-aiplatform.googleapis.com" + self._script: list[Response] = list(script) if isinstance(script, list) else [] + self._responder = script if callable(script) else None + self._aux = aux or (lambda _rec: Say("Summary of the earlier conversation (fake).")) + self.default_text = default_text + self.prompt_tokens_fn = prompt_tokens_fn + self.token_policy = TokenPolicy() + self.requests: list[dict[str, Any]] = [] + self.token_requests: list[dict[str, Any]] = [] + self.connects: list[dict[str, Any]] = [] + self.issued_signatures: list[str] = [] + self.signature_by_call: dict[str, str | None] = {} + self._tokens: dict[str, _Token] = {} + self._lock = threading.Lock() + self._seq = 0 + self._server: ThreadingHTTPServer | None = None + self._sa_project = sa_project or project + self.sa: ServiceAccount | None = None + self.ca_pem: Path | None = None + self._tls: ssl.SSLContext | None = None + + # lifecycle + def __enter__(self) -> "FakeVertex": + self.start() + return self + + def __exit__(self, *_exc: object) -> None: + self.stop() + + def start(self) -> None: + server = ThreadingHTTPServer(("127.0.0.1", 0), _handler_for(self)) + server.daemon_threads = True + self._server = server + self.ca_pem, cert, key = make_tls_material(self.root / "tls", [self.host]) + self._tls = ssl.create_default_context(ssl.Purpose.CLIENT_AUTH) + self._tls.load_cert_chain(cert, key) + self.sa = make_service_account(self.root / "sa.json", self.token_uri, self._sa_project) + threading.Thread(target=server.serve_forever, name="fake-vertex", daemon=True).start() + + def stop(self) -> None: + if self._server is not None: + self._server.shutdown() + self._server.server_close() + + @property + def port(self) -> int: + assert self._server is not None + return self._server.server_address[1] + + @property + def token_uri(self) -> str: + return f"http://127.0.0.1:{self.port}/token" + + @property + def chat_path(self) -> str: + return f"/v1beta1/projects/{self.project}/locations/{self.region}/endpoints/openapi/chat/completions" + + def child_env(self) -> dict[str, str]: + """Standard proxy + CA-trust env for the Hermes child (no Hermes-specific knobs).""" + return {"HTTPS_PROXY": f"http://127.0.0.1:{self.port}", "NO_PROXY": "127.0.0.1,localhost", + "SSL_CERT_FILE": str(self.ca_pem)} + + # scripting + def push(self, *responses: Response) -> None: + with self._lock: + self._script.extend(responses) + + def _next_main(self, record: dict[str, Any]) -> Response: + if self._responder is not None: + return self._responder(record) + with self._lock: + return self._script.pop(0) if self._script else Say(self.default_text) + + def expire_all_tokens(self) -> None: + with self._lock: + for tok in self._tokens.values(): + tok.expires_at = 0.0 + + # inspection + def main_requests(self) -> list[dict[str, Any]]: + return [r for r in self.requests if r["kind"] == "main"] + + def aux_requests(self) -> list[dict[str, Any]]: + return [r for r in self.requests if r["kind"] == "aux"] + + def rejected(self) -> list[dict[str, Any]]: + return [r for r in self.requests if r.get("rejected")] + + def minted_tokens(self) -> list[str]: + return [t["access_token"] for t in self.token_requests if t.get("access_token")] + + # token endpoint + def _mint(self, form: dict[str, str]) -> tuple[int, dict[str, Any]]: + rec: dict[str, Any] = {"grant_type": form.get("grant_type"), "t": time.time()} + with self._lock: + self.token_requests.append(rec) + policy = self.token_policy + if form.get("grant_type") != JWT_BEARER_GRANT: + rec["error"] = "unsupported_grant_type" + return 400, {"error": "unsupported_grant_type", "error_description": "Invalid grant_type"} + assert self.sa is not None + claims, why = verify_jwt_assertion(form.get("assertion", ""), self.sa) + rec["claims"] = claims + if claims is None: + rec["error"] = why + return 400, {"error": "invalid_grant", "error_description": why} + if policy.error: + rec["error"] = policy.error[1] + return policy.error[0], {"error": policy.error[1], "error_description": policy.error[2]} + value = f"ya29.fake-{secrets.token_urlsafe(24)}" + with self._lock: + self._tokens[value] = _Token(value, time.time() + policy.expires_in) + rec["access_token"] = value + return 200, {"access_token": value, "expires_in": policy.expires_in, "token_type": "Bearer"} + + def _auth_ok(self, header: str) -> bool: + if self.token_policy.reject_bearers: + return False + scheme, _, value = header.partition(" ") + tok = self._tokens.get(value) if scheme == "Bearer" else None + return bool(tok and tok.expires_at > time.time()) + + def next_ids(self) -> tuple[str, str]: + with self._lock: + self._seq += 1 + sig = base64.b64encode(f"sig-{self._seq}-".encode() + secrets.token_bytes(24)).decode() + return f"function-call-{self._seq}{secrets.randbelow(10**6):06d}", sig + + +def _vertex_error(status: int, message: str, list_body: bool = True) -> bytes: + err = {"error": {"code": status, "message": message, "status": GRPC_STATUS.get(status, "UNKNOWN")}} + return json.dumps([err] if list_body else err).encode() + + +def _chunk(model: str, delta: dict[str, Any], finish: str | None = None, usage: dict | None = None) -> dict[str, Any]: + out: dict[str, Any] = {"id": "vertex-fake", "object": "chat.completion.chunk", "created": int(time.time()), + "model": model, "choices": [{"index": 0, "delta": delta, "finish_reason": finish}]} + if usage is not None: + out["usage"] = usage + return out + + +def _handler_for(fake: FakeVertex) -> type[BaseHTTPRequestHandler]: + class Handler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + tunnel_host: str | None = None + + def log_message(self, format: str, *args: Any) -> None: # noqa: A002 + pass + + def _send(self, status: int, body: bytes, headers: dict[str, str] | None = None) -> None: + self.send_response(status) + self.send_header("Content-Type", "application/json; charset=UTF-8") + self.send_header("Content-Length", str(len(body))) + for k, v in (headers or {}).items(): + self.send_header(k, v) + self.end_headers() + self.wfile.write(body) + self.wfile.flush() + + # proxy + def do_CONNECT(self) -> None: # noqa: N802 + host, _, port = self.path.partition(":") + allowed = host == fake.host and port == "443" + with fake._lock: + fake.connects.append({"target": self.path, "allowed": allowed, "t": time.time()}) + if not allowed: + self._send(403, b'{"error": "fake proxy: host not allowed"}') + self.close_connection = True + return + self.send_response(200, "Connection Established") + self.end_headers() + self.wfile.flush() + assert fake._tls is not None + try: + tls = fake._tls.wrap_socket(self.connection, server_side=True) + except (ssl.SSLError, OSError): + self.close_connection = True + return + self.connection = tls + self.rfile = tls.makefile("rb") + self.wfile = socketserver._SocketWriter(tls) # type: ignore[attr-defined] + self.tunnel_host = host + self.close_connection = False + + # token endpoint (plain loopback) and Vertex (inside the TLS tunnel) + def do_POST(self) -> None: # noqa: N802 + raw = self.rfile.read(int(self.headers.get("Content-Length", 0) or 0)) + if self.tunnel_host is None: + if self.path != "/token": + self._send(404, b'{"error": "not found"}') + return + form = dict(urllib.parse.parse_qsl(raw.decode())) + status, payload = fake._mint(form) + self._send(status, json.dumps(payload).encode()) + return + self._vertex(raw) + + def do_GET(self) -> None: # noqa: N802 + self._send(404, _vertex_error(404, f"The requested URL {self.path} was not found on this server.")) + + def _vertex(self, raw: bytes) -> None: + record: dict[str, Any] = { + "host": self.tunnel_host, "path": self.path, "auth": self.headers.get("Authorization", ""), + "headers": {k.lower(): v for k, v in self.headers.items()}, "t": time.time(), "rejected": None, + } + try: + body = json.loads(raw or b"{}") + except json.JSONDecodeError: + body = None + record["body"] = body + record["kind"] = "main" if isinstance(body, dict) and body.get("tools") else "aux" + with fake._lock: + fake.requests.append(record) + problem = self._precheck(record) + if problem: + record["rejected"], record["status"] = problem[1], problem[0] + self._send(problem[0], _vertex_error(problem[0], problem[1])) + return + resp = fake._next_main(record) if record["kind"] == "main" else fake._aux(record) + record["response"] = type(resp).__name__ + self._respond(resp, record) + if getattr(resp, "expire_tokens_after", False): + fake.expire_all_tokens() + + def _precheck(self, record: dict[str, Any]) -> tuple[int, str] | None: + if record["host"] != fake.host or record["path"] != fake.chat_path: + return 404, f"Resource not found: {record['host']}{record['path']} (expected {fake.host}{fake.chat_path})" + if not fake._auth_ok(record["auth"]): + return 401, UNAUTHENTICATED_MESSAGE + if not isinstance(record["body"], dict): + return 400, "Invalid JSON payload received." + bad = validate_chat_body(record["body"], set(fake.issued_signatures)) + return (400, bad) if bad else None + + # rendering + def _respond(self, resp: Response, record: dict[str, Any]) -> None: + body = record["body"] + model = body.get("model", "google/gemini") + if isinstance(resp, Fail): + record["status"] = resp.status + headers = {"Retry-After": str(resp.retry_after)} if resp.retry_after is not None else {} + self._send(resp.status, _vertex_error(resp.status, resp.message, resp.list_body), headers) + return + record["status"] = 200 + if isinstance(resp, Drop): + self._start_sse() + self._sse(_chunk(model, {"role": "assistant", "content": resp.text[: resp.after_chars]})) + self.connection.close() # no terminal chunk: the chunked body is left incomplete + self.close_connection = True + return + pt = resp.prompt_tokens if resp.prompt_tokens is not None else ( + fake.prompt_tokens_fn(body) if fake.prompt_tokens_fn else 120) + deltas, finish = self._deltas(resp, record) + usage = {"prompt_tokens": pt, "completion_tokens": 20, "total_tokens": pt + 20} + if not body.get("stream"): + message: dict[str, Any] = {"role": "assistant", "content": "".join(d.get("content", "") for d in deltas) or None} + calls = [tc for d in deltas for tc in d.get("tool_calls", [])] + if calls: + message["tool_calls"] = [{k: v for k, v in tc.items() if k != "index"} for tc in calls] + self._send(200, json.dumps({"id": "vertex-fake", "object": "chat.completion", "created": int(time.time()), + "model": model, "usage": usage, + "choices": [{"index": 0, "message": message, "finish_reason": finish}]}).encode()) + return + self._start_sse() + for delta in deltas: + self._sse(_chunk(model, {"role": "assistant", **delta})) + self._sse(_chunk(model, {}, finish, usage)) + self._write_chunk(b"data: [DONE]\n\n") + self._write_chunk(b"") + self.close_connection = True + + def _deltas(self, resp: Say | Call, record: dict[str, Any]) -> tuple[list[dict[str, Any]], str]: + if isinstance(resp, Say): + size = max(1, resp.chunk_chars) + return [{"content": resp.text[i:i + size]} for i in range(0, len(resp.text), size)] or [{"content": ""}], "stop" + tool_calls = [] + for pos, (name, args) in enumerate(resp.calls): + call_id, sig = fake.next_ids() + tc: dict[str, Any] = {"index": pos, "id": call_id, "type": "function", + "function": {"name": name, "arguments": json.dumps(args)}} + if resp.signed and pos == 0: + tc["extra_content"] = {"google": {"thought_signature": sig}} + with fake._lock: + fake.issued_signatures.append(sig) + tool_calls.append(tc) + record["tool_calls"] = tool_calls + with fake._lock: + fake.signature_by_call.update({tc["id"]: _first_signature(tc) for tc in tool_calls}) + deltas: list[dict[str, Any]] = [{"content": resp.text}] if resp.text else [] + deltas.append({"tool_calls": tool_calls}) + return deltas, "tool_calls" + + def _start_sse(self) -> None: + self.send_response(200) + self.send_header("Content-Type", "text/event-stream") + self.send_header("Transfer-Encoding", "chunked") + self.end_headers() + self.wfile.flush() + + def _write_chunk(self, data: bytes) -> None: + self.wfile.write(f"{len(data):x}\r\n".encode() + data + b"\r\n") + self.wfile.flush() + + def _sse(self, payload: dict[str, Any]) -> None: + self._write_chunk(f"data: {json.dumps(payload)}\n\n".encode()) + + return Handler + + +MODEL = "google/gemini-3-flash-preview" +PROJECT = "hermes-e2e-proj" +SA_EMBEDDED_PROJECT = "hermes-sa-embedded-proj" # differs from PROJECT: proves the config override wins +REGION = "us-central1" + + +def hermes_setup(fake: FakeVertex, *, model: str = MODEL, extra_config: dict[str, Any] | None = None, + context_length: int | None = None) -> dict[str, Any]: + """``make_home`` kwargs for a Hermes home that selects ``provider: vertex`` against ``fake``: the SA + key path in ``.env`` (VERTEX_CREDENTIALS_PATH) and project/region under ``vertex:`` in config.yaml.""" + assert fake.sa is not None + block: dict[str, Any] = {"provider": "vertex", "default": model} + if context_length: + block["context_length"] = context_length + cfg: dict[str, Any] = {"vertex": {"project_id": fake.project, "region": fake.region}} + cfg.update(extra_config or {}) + return {"model": block, "env_file": {"VERTEX_CREDENTIALS_PATH": str(fake.sa.path)}, "extra_config": cfg} + + +def signatures_on_wire(body: dict[str, Any]) -> list[str]: + """Every thought signature replayed in a request's assistant tool calls, in order.""" + return [sig for m in body.get("messages", []) if m.get("role") == "assistant" + for tc in m.get("tool_calls") or [] if (sig := _first_signature(tc))] + + +__all__ = [ + "MODEL", "PROJECT", "REGION", "SA_EMBEDDED_PROJECT", "hermes_setup", + "Call", "Drop", "Fail", "FakeVertex", "GRPC_STATUS", "Say", "ServiceAccount", "TokenPolicy", + "make_service_account", "make_tls_material", "signatures_on_wire", "validate_chat_body", "verify_jwt_assertion", +]