Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Empty file.
195 changes: 195 additions & 0 deletions tests/e2e/core/providers/_native_helpers.py
Original file line number Diff line number Diff line change
@@ -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 <id>``) 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})"
173 changes: 173 additions & 0 deletions tests/e2e/core/providers/test_native_bedrock_converse_faults.py
Original file line number Diff line number Diff line change
@@ -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
Loading
Loading