158 lines
6 KiB
Python
158 lines
6 KiB
Python
"""Shared e2e fixtures (Step 13): a scripted chat client that emits a SYNTHETIC UsageDetails
|
|
(so token accounting is real-shaped without an LLM), plus store + docs-dir fixtures.
|
|
|
|
The synthetic ``UsageDetails`` is what lets the budget meter / provenance ``token_usage`` be a
|
|
positive, UsageDetails-sourced number in CI — the REAL-provider populated-usage assertion is
|
|
the gated live arm (Step 14).
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from collections.abc import Callable, Sequence
|
|
|
|
import pytest
|
|
from agent_framework import BaseChatClient
|
|
|
|
from portfolio_optimiser.simulation import ScriptedChatClient
|
|
from portfolio_optimiser.verdicts import VerdictStore, seed_store
|
|
|
|
|
|
class SyntheticUsageChatClient(ScriptedChatClient):
|
|
"""The scripted-list-then-default test double — now a THIN subclass of the canonical
|
|
``ScriptedChatClient`` (S2.5 consolidation). It keeps its full PUBLIC surface (the
|
|
``default_reply=`` kwarg constructor, the ``call_count`` attribute, ``model``/OTEL
|
|
``"synthetic"``) but delegates the shared ``_inner_get_response`` body to the canonical — the
|
|
scripted-list-then-default behaviour lives in its selector."""
|
|
|
|
def __init__(
|
|
self,
|
|
scripted: Sequence[str] | None = None,
|
|
*,
|
|
default_reply: str = "ok",
|
|
tokens_per_reply: int = 8,
|
|
) -> None:
|
|
scripted_list = list(scripted or [])
|
|
counter = {"i": 0}
|
|
|
|
def _select(_blob: str, _role: str) -> str:
|
|
i = counter["i"]
|
|
counter["i"] = i + 1
|
|
return scripted_list[i] if i < len(scripted_list) else default_reply
|
|
|
|
super().__init__(
|
|
reply_selector=_select, default_reply=default_reply, tokens_per_reply=tokens_per_reply
|
|
)
|
|
|
|
|
|
@pytest.fixture()
|
|
def make_client_factory() -> Callable[..., Callable[[str], BaseChatClient]]:
|
|
"""Return a maker that builds a per-role client factory emitting synthetic usage."""
|
|
|
|
def _make(default_reply: str, *, tokens: int = 8) -> Callable[[str], BaseChatClient]:
|
|
def factory(role: str) -> BaseChatClient:
|
|
return SyntheticUsageChatClient(default_reply=default_reply, tokens_per_reply=tokens)
|
|
|
|
return factory
|
|
|
|
return _make
|
|
|
|
|
|
# A generic VALID SavingsProposal reply for any project not present in a portfolio reply map:
|
|
# affected total = 1 x 100_000 = 100_000, P90 = 0.30 x 100_000 = 30_000, claimed 20_000 <= both
|
|
# (Pydantic affected-total invariant and the validator P90 gate) -> always validates.
|
|
_PORTFOLIO_DEFAULT_REPLY = (
|
|
'{"measure":"Reduce scope","affected_items":'
|
|
'[{"code":"01.1","quantity":1,"unit_cost":100000}],"claimed_saving_nok":20000}'
|
|
)
|
|
|
|
|
|
class _ProjectAwareUsageChatClient(ScriptedChatClient):
|
|
"""Selects its reply by scanning the incoming prompt for a known ``project_id`` substring (the
|
|
prompt embeds ``project.id`` at run.py:162 and generate.py:48), falling back to a default valid
|
|
proposal — so ``run_portfolio``'s single ``client_factory`` stays production-shaped while tests
|
|
vary the proposal per project. A THIN subclass: the prompt-scan lives in its selector, the shared
|
|
``_inner_get_response`` body in the canonical."""
|
|
|
|
def __init__(
|
|
self, replies: dict[str, str], *, default_reply: str, tokens_per_reply: int = 8
|
|
) -> None:
|
|
table = dict(replies)
|
|
|
|
def _select(blob: str, _role: str) -> str:
|
|
return next((r for pid, r in table.items() if pid in blob), default_reply)
|
|
|
|
super().__init__(
|
|
reply_selector=_select, default_reply=default_reply, tokens_per_reply=tokens_per_reply
|
|
)
|
|
|
|
|
|
@pytest.fixture()
|
|
def make_portfolio_client_factory() -> Callable[..., Callable[[str], BaseChatClient]]:
|
|
"""Return a maker that builds a single project-aware client factory: every client it
|
|
produces picks its reply from ``replies`` by scanning the prompt for the project id, so one
|
|
factory serves the whole portfolio (matching ``run_portfolio``'s single-factory seam)."""
|
|
|
|
def _make(
|
|
replies: dict[str, str],
|
|
*,
|
|
default_reply: str = _PORTFOLIO_DEFAULT_REPLY,
|
|
tokens: int = 8,
|
|
) -> Callable[[str], BaseChatClient]:
|
|
def factory(role: str) -> BaseChatClient:
|
|
return _ProjectAwareUsageChatClient(
|
|
replies, default_reply=default_reply, tokens_per_reply=tokens
|
|
)
|
|
|
|
return factory
|
|
|
|
return _make
|
|
|
|
|
|
class _RecordingChatClient(ScriptedChatClient):
|
|
"""Records the incoming prompt blob per call into a SHARED sink, then returns a fixed valid
|
|
reply. Lets a test assert exactly what text reached the prompt — the probe the Step-1 ExpeL
|
|
wiring is made load-bearing against (does a prior verdict reach the hypothesis prompt?). A THIN
|
|
subclass: the canonical records to the ``sink`` (when given one) and returns the constant reply."""
|
|
|
|
def __init__(self, sink: list[str], reply: str, *, tokens_per_reply: int = 8) -> None:
|
|
super().__init__(reply, sink, tokens_per_reply=tokens_per_reply)
|
|
|
|
|
|
@pytest.fixture()
|
|
def make_recording_client_factory() -> Callable[
|
|
[str], tuple[Callable[[str], BaseChatClient], list[str]]
|
|
]:
|
|
"""Return a maker that builds a per-role client factory recording every prompt blob into a
|
|
shared list. Returns ``(factory, recorded_prompts)`` so the test inspects what reached the
|
|
prompt across the whole run (debate rounds + generation)."""
|
|
|
|
def _make(reply: str) -> tuple[Callable[[str], BaseChatClient], list[str]]:
|
|
sink: list[str] = []
|
|
|
|
def factory(role: str) -> BaseChatClient:
|
|
return _RecordingChatClient(sink, reply)
|
|
|
|
return factory, sink
|
|
|
|
return _make
|
|
|
|
|
|
@pytest.fixture()
|
|
def fresh_store() -> VerdictStore:
|
|
return VerdictStore(verdicts=[])
|
|
|
|
|
|
@pytest.fixture()
|
|
def seeded_store() -> VerdictStore:
|
|
return seed_store()
|
|
|
|
|
|
@pytest.fixture()
|
|
def docs_dir(tmp_path) -> str:
|
|
d = tmp_path / "docs"
|
|
d.mkdir()
|
|
(d / "cost.txt").write_text(
|
|
"Asphalt Ab11 unit rate renegotiation reduced the paving cost on the school stretch.",
|
|
encoding="utf-8",
|
|
)
|
|
return str(d)
|