diff --git a/reflexio/models/api_schema/domain/entities.py b/reflexio/models/api_schema/domain/entities.py index dceec3d69..a90a34ba1 100644 --- a/reflexio/models/api_schema/domain/entities.py +++ b/reflexio/models/api_schema/domain/entities.py @@ -507,6 +507,11 @@ class LineageEvent(BaseModel): request_id (str): Triggering request — part of the idempotency key. reason (str): Free-text rationale (no PII). created_at (int): Unix epoch seconds (0 = unset; storage stamps it). + from_status (str | None): Status before a transition. + to_status (str | None): Status after a transition. + status_namespace (str | None): Namespace for status values. + model_name (str | None): Observed model for a content-shaping operation. + provider (str | None): Observed provider for that operation. """ event_id: int = 0 @@ -523,6 +528,8 @@ class LineageEvent(BaseModel): from_status: str | None = None to_status: str | None = None status_namespace: str | None = None + model_name: str | None = None + provider: str | None = None class LineageContext(BaseModel): @@ -536,6 +543,8 @@ class LineageContext(BaseModel): source_ids: list[str] = [] reason: str = "" request_id: str | None = None + model_name: str | None = None + provider: str | None = None class RecordRef(BaseModel): diff --git a/reflexio/server/llm/_litellm_text_generation.py b/reflexio/server/llm/_litellm_text_generation.py index a561297a5..1ee3c7e94 100644 --- a/reflexio/server/llm/_litellm_text_generation.py +++ b/reflexio/server/llm/_litellm_text_generation.py @@ -42,8 +42,10 @@ from reflexio.server.llm._litellm_subprocess import _litellm_completion_worker from reflexio.server.llm._litellm_types import ( + CompletionResult, LiteLLMClientError, LLMHardTimeoutError, + ModelProvenance, StructuredOutputParseError, StructuredOutputRepairError, ToolCallingChatResponse, @@ -122,6 +124,19 @@ StructuredOutputValidator = Callable[[BaseModel], Sequence[str]] +def _nonempty_string(value: Any) -> str | None: + """Return trustworthy string metadata without coercing mocks or objects.""" + if not isinstance(value, str): + return None + value = value.strip() + return value or None + + +def _response_hidden_params(response: Any) -> dict[str, Any]: + hidden = getattr(response, "_hidden_params", None) + return hidden if isinstance(hidden, dict) else {} + + @dataclass class _StructuredAttempt: value: str | BaseModel | ToolCallingChatResponse @@ -129,6 +144,7 @@ class _StructuredAttempt: parsed_output: BaseModel | None finish_reason: str | None model: str + provenance: ModelProvenance | None = None def _is_expected_transient_llm_error(exc: BaseException) -> bool: @@ -247,7 +263,23 @@ def generate_response( LiteLLMClientError: If the API call fails after all retries, or if response_format is not a Pydantic BaseModel class. """ - # Validate response_format if provided + return self.generate_response_with_provenance( + prompt, + system_message, + images, + image_media_type, + **kwargs, + ).value + + def generate_response_with_provenance( + self, + prompt: str, + system_message: str | None = None, + images: list[str | bytes | dict] | None = None, + image_media_type: str | None = None, + **kwargs: Any, + ) -> CompletionResult[str | BaseModel | ToolCallingChatResponse]: + """Generate a response paired with observed model provenance.""" response_format = kwargs.get("response_format") if response_format is not None and not is_pydantic_model(response_format): raise LiteLLMClientError( @@ -316,7 +348,32 @@ def generate_chat_response( LiteLLMClientError: If the API call fails after all retries, or if response_format is not a Pydantic BaseModel class. """ - # Validate response_format if provided + return self.generate_chat_response_with_provenance( + messages, + system_message, + tools=tools, + tool_choice=tool_choice, + model_role=model_role, + max_retries=max_retries, + fallback_models=fallback_models, + structured_output_validator=structured_output_validator, + **kwargs, + ).value + + def generate_chat_response_with_provenance( + self, + messages: list[dict[str, Any]], + system_message: str | None = None, + *, + tools: list[Any] | None = None, + tool_choice: str | dict[str, Any] | None = None, + model_role: ModelRole | None = None, + max_retries: int | None = None, + fallback_models: list[str] | None = None, + structured_output_validator: StructuredOutputValidator | None = None, + **kwargs: Any, + ) -> CompletionResult[str | BaseModel | ToolCallingChatResponse]: + """Generate a chat response paired with observed model provenance.""" response_format = kwargs.get("response_format") if response_format is not None and not is_pydantic_model(response_format): raise LiteLLMClientError( @@ -866,6 +923,46 @@ def _log_token_usage(self, params: dict[str, Any], response: Any) -> None: cost_suffix, ) + def _build_model_provenance(self, response: Any) -> ModelProvenance: + """Build attribution only from metadata observed on the response. + + Model name is taken only from fields that represent the served completion, + not request-side LiteLLM routing metadata. ``_hidden_params["model"]`` is + intentionally ignored: LiteLLM often echoes the requested model there, which + would record a configured route as if it had been observed. + + The claude-code LiteLLM route can execute different local host CLIs. Its + public ``ModelResponse.model`` is the *requested* route string (same as + other providers). Observed model is only the bridge stamp + ``reflexio_served_model`` — never ``response.model``, which would launder + the requested route as observed when the CLI does not report a served model. + """ + hidden = _response_hidden_params(response) + route_provider = _nonempty_string(hidden.get("reflexio_provider")) + served_provider = _nonempty_string(hidden.get("reflexio_served_provider")) + stamped_served = _nonempty_string(hidden.get("reflexio_served_model")) + + if route_provider == "claude-code": + provider = served_provider + model_name = stamped_served + else: + # Prefer an explicit bridge stamp, then the provider response body field. + # Do not fall back to request-side hidden model / model_id keys. + model_name = stamped_served or _nonempty_string( + getattr(response, "model", None) + ) + provider = ( + served_provider + or _nonempty_string(hidden.get("custom_llm_provider")) + or _nonempty_string(hidden.get("llm_provider")) + or _nonempty_string(hidden.get("provider")) + ) + + return ModelProvenance( + model_name=model_name, + provider=provider, + ) + def _emit_fallback_signal( self, primary_model: str, served_model: str, *, reason: str ) -> None: @@ -907,7 +1004,7 @@ def _emit_fallback_signal( def _make_request( # noqa: C901 self, messages: list[dict[str, Any]], **kwargs: Any - ) -> str | BaseModel | ToolCallingChatResponse: + ) -> CompletionResult[str | BaseModel | ToolCallingChatResponse]: """ Make a request to the LLM via a reflexio-owned per-rung fallback walk. @@ -943,6 +1040,12 @@ def _make_request( # noqa: C901 ) original_kwargs = dict(kwargs) + def _finish( + attempt: _StructuredAttempt, + ) -> CompletionResult[str | BaseModel | ToolCallingChatResponse]: + assert attempt.provenance is not None # noqa: S101 + return CompletionResult(attempt.value, attempt.provenance) + if structured_output_validator is not None and ( original_kwargs.get("response_format") is None or not original_kwargs.get("parse_structured_output", True) @@ -1007,6 +1110,7 @@ def _call_and_parse( response = self._completion_with_hard_timeout( turn_params, turn_hard_timeout ) + provenance = self._build_model_provenance(response) message = response.choices[0].message # type: ignore[reportAttributeAccessIssue] content = message.content finish_reason = response.choices[0].finish_reason # type: ignore[reportAttributeAccessIssue] @@ -1046,6 +1150,7 @@ def _call_and_parse( exc.finish_reason = finish_reason if exc.raw_content is None and isinstance(content, str): exc.raw_content = content + exc.provenance = provenance raise if isinstance(parsed, BaseModel): parsed_output = parsed @@ -1063,6 +1168,7 @@ def _call_and_parse( parsed_output=parsed_output, finish_reason=finish_reason, model=str(turn_params.get("model")), + provenance=provenance, ) try: @@ -1075,6 +1181,7 @@ def _call_and_parse( exc.finish_reason = finish_reason if exc.raw_content is None and isinstance(content, str): exc.raw_content = content + exc.provenance = provenance raise return _StructuredAttempt( value=value, @@ -1082,6 +1189,7 @@ def _call_and_parse( parsed_output=value if isinstance(value, BaseModel) else None, finish_reason=finish_reason, model=str(turn_params.get("model")), + provenance=provenance, ) except ( StructuredOutputParseError, @@ -1197,6 +1305,7 @@ def _repair_error( attempt: _StructuredAttempt | None, errors: Sequence[str], model: str, + first_parsed_provenance: ModelProvenance | None = None, ) -> StructuredOutputRepairError: return StructuredOutputRepairError( "Structured output repair exhausted", @@ -1205,11 +1314,12 @@ def _repair_error( raw_content=attempt.raw_content if attempt else None, parsed_output=attempt.parsed_output if attempt else None, validation_errors=tuple(errors), + first_parsed_provenance=first_parsed_provenance, ) def _run_rung_plain( rung_messages: list[dict[str, Any]], rung_kwargs: dict[str, Any] - ) -> str | BaseModel | ToolCallingChatResponse: + ) -> CompletionResult[str | BaseModel | ToolCallingChatResponse]: """Serve one rung with no validator: initial call + one same-model parse-retry. Raises ``StructuredOutputParseError`` when both attempts return a @@ -1221,22 +1331,26 @@ def _run_rung_plain( rung_messages, rung_kwargs ) try: - return _call_and_parse( - params, rf, parse_so, hard_timeout, detect_refusal=False - ).value + return _finish( + _call_and_parse( + params, rf, parse_so, hard_timeout, detect_refusal=False + ) + ) except StructuredOutputParseError: self.logger.warning( "event=llm_parse_retry model=%s — malformed structured output, " "retrying once on the same model", params.get("model"), ) - return _call_and_parse( - params, rf, parse_so, hard_timeout, detect_refusal=False - ).value + return _finish( + _call_and_parse( + params, rf, parse_so, hard_timeout, detect_refusal=False + ) + ) def _run_rung_validated( rung_messages: list[dict[str, Any]], rung_kwargs: dict[str, Any] - ) -> str | BaseModel | ToolCallingChatResponse: + ) -> CompletionResult[str | BaseModel | ToolCallingChatResponse]: """Serve one rung with the validator: initial call + one same-model repair turn. The corrective turn is built from ``rung_messages`` (this rung's @@ -1249,6 +1363,7 @@ def _run_rung_validated( ) schema_name = getattr(rf, "__name__", "structured output") latest_parsed_output: BaseModel | None + first_parsed_provenance: ModelProvenance | None = None try: first_attempt = _call_and_parse( params, rf, parse_so, hard_timeout, detect_refusal=True @@ -1262,10 +1377,12 @@ def _run_rung_validated( else: valid, errors, failure_kind = _validate_attempt(first_attempt) if valid: - return first_attempt.value + return _finish(first_attempt) raw_content = first_attempt.raw_content finish_reason = first_attempt.finish_reason latest_parsed_output = first_attempt.parsed_output + if first_attempt.parsed_output is not None: + first_parsed_provenance = first_attempt.provenance repair_base = _repair_messages( rung_messages, @@ -1300,9 +1417,15 @@ def _run_rung_validated( parsed_output=None, finish_reason=exc.finish_reason, model=str(repair_params.get("model")), + provenance=getattr(exc, "provenance", None), ) errors = (str(exc),) failure_kind = "parse" + except (LiteLLMClientError, ProviderCapSaturatedError) as exc: + # Cap saturation is not a LiteLLMClientError subclass; still stamp + # first-parsed so the ladder outer handler can salvage attribution. + exc.first_parsed_provenance = first_parsed_provenance + raise else: valid, errors, failure_kind = _validate_attempt(repair_attempt) if valid: @@ -1312,7 +1435,12 @@ def _run_rung_validated( repair_params.get("model"), schema_name, ) - return repair_attempt.value + return _finish(repair_attempt) + if ( + repair_attempt.parsed_output is not None + and first_parsed_provenance is None + ): + first_parsed_provenance = repair_attempt.provenance # Within-rung roll-forward: keep the most recent output that parsed # (e.g. a semantic-fail before a final parse-fail) on the typed error. @@ -1330,6 +1458,7 @@ def _run_rung_validated( attempt=repair_attempt, errors=errors, model=repair_attempt.model, + first_parsed_provenance=first_parsed_provenance, ) # Reflexio-owned per-rung walk. Each rung is entered at most once; the @@ -1350,6 +1479,10 @@ def _run_rung_validated( ) last_error: Exception | None = None + # First accepted parse across the whole walk — not per-rung. Consolidator + # salvage keeps the first parsed *content* via a shared validator closure; + # this field must match that content's served model, not the last rung's. + ladder_first_parsed_provenance: ModelProvenance | None = None for index, rung in enumerate(ladder): rung_kwargs = {**original_kwargs, "model": rung, "fallback_models": []} # ``model_role`` is already resolved into ``ladder``; leaving it in @@ -1367,14 +1500,25 @@ def _run_rung_validated( ProviderCapSaturatedError, StructuredOutputParseError, ) as exc: + first_parsed = ( + exc.provenance + if isinstance(exc, StructuredOutputParseError) + else exc.first_parsed_provenance + ) + if first_parsed is not None and ladder_first_parsed_provenance is None: + ladder_first_parsed_provenance = first_parsed last_error = exc if not is_last: continue # Final rung failed. Preserve the typed repair error (callers keep # the latest parse) and already-wrapped client errors as-is; wrap a # raw plain-path parse exhaustion (litellm saw a 200, so no turn - # logged a request-end failure) and a cap-saturation. + # logged a request-end failure) and a cap-saturation. Always stamp + # ladder-wide first-parsed so consolidator salvage keeps matching + # model attribution. if isinstance(exc, StructuredOutputRepairError | LiteLLMClientError): + if ladder_first_parsed_provenance is not None: + exc.first_parsed_provenance = ladder_first_parsed_provenance raise if isinstance(exc, StructuredOutputParseError): self.logger.error( @@ -1383,7 +1527,11 @@ def _run_rung_validated( type(exc).__name__, exc, ) - raise LiteLLMClientError(f"API call failed: {exc}") from exc + wrapped = LiteLLMClientError( + f"API call failed: {exc}", + first_parsed_provenance=ladder_first_parsed_provenance, + ) + raise wrapped from exc else: if index > 0: self._emit_fallback_signal( @@ -1393,7 +1541,8 @@ def _run_rung_validated( # A non-empty ladder always returns or raises above; guard the empty case. raise LiteLLMClientError( # pragma: no cover - f"All fallback rungs failed; last: {last_error}" + f"All fallback rungs failed; last: {last_error}", + first_parsed_provenance=ladder_first_parsed_provenance, ) def _apply_prompt_caching( diff --git a/reflexio/server/llm/_litellm_types.py b/reflexio/server/llm/_litellm_types.py index 9ccfdc691..c8afff553 100644 --- a/reflexio/server/llm/_litellm_types.py +++ b/reflexio/server/llm/_litellm_types.py @@ -19,6 +19,22 @@ from reflexio.models.config_schema import APIKeyConfig +@dataclass(frozen=True) +class ModelProvenance: + """Observed model and provider attribution for one completion.""" + + model_name: str | None = None + provider: str | None = None + + +@dataclass(frozen=True) +class CompletionResult[T]: + """Completion value paired with its non-serializing provenance.""" + + value: T + provenance: ModelProvenance + + @dataclass class LiteLLMConfig: """ @@ -104,7 +120,20 @@ class ToolCallingChatResponse: class LiteLLMClientError(Exception): - """Custom exception for LiteLLM client errors.""" + """Custom exception for LiteLLM client errors. + + ``first_parsed_provenance`` is populated when a later structured-output + repair transport failure leaves a parsed response available to a caller. + """ + + def __init__( + self, + message: str, + *, + first_parsed_provenance: ModelProvenance | None = None, + ) -> None: + super().__init__(message) + self.first_parsed_provenance = first_parsed_provenance class StructuredOutputRepairError(LiteLLMClientError): @@ -112,9 +141,10 @@ class StructuredOutputRepairError(LiteLLMClientError): Field pairing caveat: ``raw_content``/``validation_errors`` describe the LAST attempt, while ``parsed_output`` falls back to the most recent attempt - that parsed at all — when the final attempt failed to parse, these fields - describe different attempts. Callers must not assume ``validation_errors`` - were produced by validating ``parsed_output``. + that parsed at all. ``first_parsed_provenance`` is the first parse across the + whole multi-rung walk (not merely the final rung), so salvage callers can + pair it with the first accepted parsed content from a shared validator + closure. """ def __init__( @@ -126,8 +156,9 @@ def __init__( raw_content: str | None = None, parsed_output: BaseModel | None = None, validation_errors: tuple[str, ...] = (), + first_parsed_provenance: ModelProvenance | None = None, ) -> None: - super().__init__(message) + super().__init__(message, first_parsed_provenance=first_parsed_provenance) self.failure_kind = failure_kind self.model = model self.raw_content = raw_content @@ -148,10 +179,12 @@ def __init__( *, raw_content: str | None = None, finish_reason: str | None = None, + provenance: ModelProvenance | None = None, ) -> None: super().__init__(message) self.raw_content = raw_content self.finish_reason = finish_reason + self.provenance = provenance class LLMHardTimeoutError(TimeoutError): diff --git a/reflexio/server/llm/_provider_concurrency.py b/reflexio/server/llm/_provider_concurrency.py index f1c452632..51ccac9c2 100644 --- a/reflexio/server/llm/_provider_concurrency.py +++ b/reflexio/server/llm/_provider_concurrency.py @@ -16,12 +16,16 @@ import threading from collections.abc import Iterator from contextlib import contextmanager +from typing import TYPE_CHECKING import litellm from reflexio.server.env_utils import env_str from reflexio.server.llm.llm_utils import positive_int_env +if TYPE_CHECKING: + from reflexio.server.llm._litellm_types import ModelProvenance + logger = logging.getLogger(__name__) _DEFAULT_MAX_CONCURRENCY = 8 @@ -44,6 +48,15 @@ class ProviderCapSaturatedError(Exception): advance-worthy rung failure. """ + def __init__( + self, + message: str, + *, + first_parsed_provenance: "ModelProvenance | None" = None, + ) -> None: + super().__init__(message) + self.first_parsed_provenance = first_parsed_provenance + def _parse_fail_closed() -> frozenset[str]: raw = env_str("REFLEXIO_LLM_FAIL_CLOSED_PROVIDERS", "") diff --git a/reflexio/server/llm/providers/claude_code_provider.py b/reflexio/server/llm/providers/claude_code_provider.py index a26a2f466..dd6356d2f 100644 --- a/reflexio/server/llm/providers/claude_code_provider.py +++ b/reflexio/server/llm/providers/claude_code_provider.py @@ -28,6 +28,7 @@ import tempfile import time from contextlib import suppress +from dataclasses import replace from datetime import UTC, datetime from pathlib import Path from typing import Any @@ -530,8 +531,11 @@ def _run_claude_stream( except FileNotFoundError as exc: raise ClaudeCodeCLIError(f"claude CLI not found at {cli_path}") from exc - return parse_stream_json( - proc.stdout, exit_code=proc.returncode, stderr_text=proc.stderr + return replace( + parse_stream_json( + proc.stdout, exit_code=proc.returncode, stderr_text=proc.stderr + ), + cli_binary=_cli_name(), ) @@ -592,6 +596,7 @@ def _run_codex_stream( terminal_text=terminal_text, stderr_text=proc.stderr, raw_lines_parsed=1 if terminal_text else 0, + cli_binary="codex", ) @@ -612,6 +617,10 @@ def _build_model_response( model: str, terminal_text: str, elapsed_seconds: float, + *, + served_model: str | None = None, + served_provider: str | None = None, + cli_binary: str | None = None, ) -> ModelResponse: """Wrap the CLI's terminal text in a LiteLLM ``ModelResponse``. @@ -621,9 +630,12 @@ def _build_model_response( Args: model (str): The model string originally requested - (e.g. ``claude-code/default``). + (e.g. ``claude-code/default``). Populates the public LiteLLM + ``ModelResponse.model`` field the same way other providers do. terminal_text (str): The terminal ``result`` text from the CLI. elapsed_seconds (float): Wall time the subprocess took — for logging only. + served_model: Observed served model from stream-json, if any. Stamped on + hidden metadata for provenance; not used to overwrite ``model``. Returns: ModelResponse: Shaped to match what callers of ``litellm.completion`` expect. @@ -639,6 +651,12 @@ def _build_model_response( object="chat.completion", usage=usage, ) + _set_cli_response_metadata( + response, + served_model=served_model, + served_provider=served_provider, + cli_binary=cli_binary, + ) _LOGGER.debug( "claude-code provider: model=%s elapsed=%.2fs", model, @@ -647,6 +665,25 @@ def _build_model_response( return response +def _set_cli_response_metadata( + response: ModelResponse, + *, + served_model: str | None, + served_provider: str | None, + cli_binary: str | None, +) -> None: + """Stamp truthful route metadata on a CLI-backed completion response.""" + hidden = dict(getattr(response, "_hidden_params", {}) or {}) + hidden["reflexio_provider"] = PROVIDER_KEY + if cli_binary: + hidden["reflexio_cli_binary"] = cli_binary + if served_model: + hidden["reflexio_served_model"] = served_model + if served_provider: + hidden["reflexio_served_provider"] = served_provider + response._hidden_params = hidden + + _TOOL_USE_INSTRUCTION_TEMPLATE = ( "## EXTERNAL TOOL-CALLING MODE\n" "\n" @@ -785,6 +822,9 @@ def _build_model_response_with_tool_call( terminal_text: str, elapsed_seconds: float, tool_use: dict[str, Any], + served_model: str | None = None, + served_provider: str | None = None, + cli_binary: str | None = None, ) -> ModelResponse: """Wrap the CLI terminal text as a ``ModelResponse`` carrying one ``tool_calls`` entry. @@ -794,7 +834,6 @@ def _build_model_response_with_tool_call( this — usage is informational, not load-bearing. Args: - model: Model string passed in by LiteLLM. terminal_text: The terminal ``result`` text from the CLI (retained for signature parity with the plain-text branch; surfaced via logging only). @@ -826,6 +865,12 @@ def _build_model_response_with_tool_call( object="chat.completion", usage=usage, ) + _set_cli_response_metadata( + response, + served_model=served_model, + served_provider=served_provider, + cli_binary=cli_binary, + ) _LOGGER.debug( "claude-code provider: tool_call name=%s elapsed=%.2fs", tool_use["name"], @@ -1018,6 +1063,9 @@ def completion( # type: ignore[override] terminal_text=result.terminal_text, elapsed_seconds=elapsed, tool_use=tool_use, + served_model=result.served_model, + served_provider=result.served_provider, + cli_binary=result.cli_binary, ) # Log a metadata-only warning (no raw payload) — the model # output can carry user content / source code; deferring the @@ -1038,6 +1086,9 @@ def completion( # type: ignore[override] model=model, terminal_text=result.terminal_text, elapsed_seconds=elapsed, + served_model=result.served_model, + served_provider=result.served_provider, + cli_binary=result.cli_binary, ) self._record_stall_safely(result) diff --git a/reflexio/server/llm/providers/claude_code_stream_parser.py b/reflexio/server/llm/providers/claude_code_stream_parser.py index 4f957d732..cce7da30f 100644 --- a/reflexio/server/llm/providers/claude_code_stream_parser.py +++ b/reflexio/server/llm/providers/claude_code_stream_parser.py @@ -63,6 +63,9 @@ class ParseResult: stderr_text: str = "" raw_lines_parsed: int = 0 raw_lines_failed: int = 0 + served_model: str | None = None + served_provider: str | None = None + cli_binary: str | None = None @property def stall_candidate(self) -> str | None: @@ -95,6 +98,12 @@ def parse_stream_json( parsed = 0 failed = 0 saw_terminal = False + init_model: str | None = None + assistant_model: str | None = None + assistant_provider: str | None = None + usage_model: str | None = None + result_model: str | None = None + result_provider: str | None = None for line in stdout.splitlines(): if not line.strip(): continue @@ -107,6 +116,10 @@ def parse_stream_json( if not isinstance(event, dict): continue match event.get("type"), event.get("subtype"): + case ("system", "init"): + model = event.get("model") + if isinstance(model, str) and model.strip(): + init_model = model case ("system", "api_retry"): err = event.get("error") if isinstance(err, str): @@ -116,6 +129,27 @@ def parse_stream_json( if isinstance(text, str): terminal_text = text saw_terminal = True + model_usage = event.get("modelUsage") + if isinstance(model_usage, dict) and len(model_usage) == 1: + model = next(iter(model_usage)) + if isinstance(model, str) and model.strip(): + usage_model = model + model = event.get("model") + if isinstance(model, str) and model.strip(): + result_model = model + provider = event.get("provider") + if isinstance(provider, str) and provider.strip(): + result_provider = provider + case ("assistant", _): + message = event.get("message") + model = message.get("model") if isinstance(message, dict) else None + if isinstance(model, str) and model.strip(): + assistant_model = model + provider = ( + message.get("provider") if isinstance(message, dict) else None + ) + if isinstance(provider, str) and provider.strip(): + assistant_provider = provider return ParseResult( success=(exit_code == 0 and saw_terminal and bool(terminal_text)), terminal_text=terminal_text, @@ -123,6 +157,8 @@ def parse_stream_json( stderr_text=stderr_text, raw_lines_parsed=parsed, raw_lines_failed=failed, + served_model=result_model or assistant_model or init_model or usage_model, + served_provider=result_provider or assistant_provider, ) diff --git a/reflexio/server/llm/tools.py b/reflexio/server/llm/tools.py index f539854b2..2fb56be10 100644 --- a/reflexio/server/llm/tools.py +++ b/reflexio/server/llm/tools.py @@ -13,6 +13,7 @@ from pydantic import BaseModel, ConfigDict, Field, ValidationError +from reflexio.server.llm._litellm_types import ModelProvenance from reflexio.server.llm.llm_utils import ( assert_provider_safe_schema, make_strict_json_schema, @@ -189,6 +190,9 @@ class ToolLoopResult(BaseModel): # (the structured-output terminus, used by the extraction agent instead of a # finish-sentinel tool call). structured_output: BaseModel | None = None + # Runtime-only attribution for the final accepted LLM turn. Persistence is + # explicit at the lineage boundary; this must not leak into API payloads. + provenance: ModelProvenance | None = Field(default=None, exclude=True) # Models we know support function calling per vendor docs but that litellm's @@ -378,16 +382,19 @@ def _run_multi_stage_fallback( log_model_response, ) + latest_provenance: ModelProvenance | None = None for turn_idx in range(max_steps): turn_label = f"(multi-stage turn {turn_idx + 1})" if log_label: log_llm_messages(logger, f"{log_label} {turn_label}", messages) tool_t0 = time.monotonic() - parsed = client.generate_chat_response( + completion = client.generate_chat_response_with_provenance( messages=messages, response_format=multi_stage_schema, model_role=model_role, ) + parsed = completion.value + latest_provenance = completion.provenance if log_label: log_model_response(logger, f"{log_label} {turn_label}", parsed) if not isinstance(parsed, BaseModel): @@ -444,6 +451,7 @@ def _run_multi_stage_fallback( messages=messages, pending_tool_call_ids=pending_tool_call_ids, max_steps_remaining=max_steps - turn_idx - 1, + provenance=latest_provenance, ) outcome = registry.handle_outcome(tool_name, args_json, ctx) @@ -474,6 +482,7 @@ def _run_multi_stage_fallback( messages=messages, pending_tool_call_ids=pending_tool_call_ids, max_steps_remaining=0, + provenance=latest_provenance, ) @@ -596,11 +605,13 @@ def run_tool_loop( ) if log_label: log_llm_messages(logger, f"{log_label} (fallback)", messages) - parsed = client.generate_chat_response( + completion = client.generate_chat_response_with_provenance( messages=messages, response_format=fallback_schema, model_role=model_role, ) + parsed = completion.value + provenance = completion.provenance if log_label: log_model_response(logger, f"{log_label} (fallback)", parsed) # The fallback path always passes response_format so the client @@ -642,6 +653,7 @@ def run_tool_loop( messages=messages, pending_tool_call_ids=pending_tool_call_ids, max_steps_remaining=0 if exceeded else max_steps - len(bounded_items), + provenance=provenance, ) # ---- Native tool loop --------------------------------------------- @@ -649,18 +661,21 @@ def run_tool_loop( from reflexio.server.llm.litellm_client import LiteLLMClientError local_msgs = list(messages) + provenance: ModelProvenance | None = None try: tool_specs = registry.openai_specs() for _step in range(max_steps): if log_label: log_llm_messages(logger, f"{log_label} (turn {_step + 1})", local_msgs) - resp = client.generate_chat_response( + completion = client.generate_chat_response_with_provenance( messages=local_msgs, tools=tool_specs or None, tool_choice=tool_choice if tool_specs else None, model_role=model_role, response_format=response_format, ) + resp = completion.value + provenance = completion.provenance if log_label: log_model_response(logger, f"{log_label} (turn {_step + 1})", resp) @@ -704,6 +719,7 @@ def run_tool_loop( # The structured answer is committed on this turn — # one LLM call consumed, mirroring the finish_tool path. max_steps_remaining=max_steps - _step - 1, + provenance=provenance, ) # No response_format requested (or nothing parseable): the finish # handler did NOT run, so no structured output was committed. @@ -719,6 +735,7 @@ def run_tool_loop( messages=local_msgs, pending_tool_call_ids=pending_tool_call_ids, max_steps_remaining=max_steps - _step, + provenance=provenance, ) normalized_tool_calls = [ _normalize_tool_call_for_history(tc) for tc in tool_calls @@ -782,6 +799,7 @@ def run_tool_loop( messages=local_msgs, pending_tool_call_ids=pending_tool_call_ids, max_steps_remaining=max_steps - _step - 1, + provenance=provenance, ) except LiteLLMClientError as e: # LLM failure after the client exhausted its retries and fallbacks — @@ -816,4 +834,5 @@ def run_tool_loop( messages=local_msgs, pending_tool_call_ids=pending_tool_call_ids, max_steps_remaining=0, + provenance=provenance, ) diff --git a/reflexio/server/services/base_generation/_extraction_lifecycle.py b/reflexio/server/services/base_generation/_extraction_lifecycle.py index 6854554f4..c27308a6d 100644 --- a/reflexio/server/services/base_generation/_extraction_lifecycle.py +++ b/reflexio/server/services/base_generation/_extraction_lifecycle.py @@ -30,6 +30,7 @@ from typing import TYPE_CHECKING, Any, Generic, TypeVar from reflexio.server.api_endpoints.request_context import RequestContext +from reflexio.server.llm._litellm_types import ModelProvenance from reflexio.server.llm.token_accounting import RunTokenTotals from reflexio.server.services.deferred_learning_plan import ExtractorBookmarkAdvance from reflexio.server.services.extraction.outcome import ExtractionOutcome @@ -63,6 +64,7 @@ class ExtractionRunLifecycleMixin(Generic[TExtractorConfig, TGenerationServiceCo _last_extraction_run_ids: list[str] _last_token_totals: RunTokenTotals | None _last_bookmark_advance: ExtractorBookmarkAdvance | None + _last_model_provenance: ModelProvenance | None if TYPE_CHECKING: # Abstract on the base ABC (stays there per SINK-2); declared here type-only so @@ -123,6 +125,7 @@ def _execute_extractor( # later in ``persist_generation`` (durable fence) or in # ``_run_generation``'s persist half for the synchronous path. self._last_bookmark_advance = result.bookmark_advance + self._last_model_provenance = result.model_provenance if result.status == "completed" and result.items: return result.items logger.info( diff --git a/reflexio/server/services/base_generation_service.py b/reflexio/server/services/base_generation_service.py index 092607ca9..a504d8dc3 100644 --- a/reflexio/server/services/base_generation_service.py +++ b/reflexio/server/services/base_generation_service.py @@ -13,6 +13,7 @@ from reflexio.models.api_schema.internal_schema import RequestInteractionDataModel from reflexio.server.api_endpoints.request_context import RequestContext +from reflexio.server.llm._litellm_types import ModelProvenance from reflexio.server.llm.litellm_client import LiteLLMClient from reflexio.server.llm.token_accounting import RunTokenTotals from reflexio.server.services.base_generation import ( @@ -257,6 +258,7 @@ def __init__( # Stride-bookmark advance deferred off the extractor (F1); captured in # ``_execute_extractor`` and applied in the persist half of the run. self._last_bookmark_advance: ExtractorBookmarkAdvance | None = None + self._last_model_provenance: ModelProvenance | None = None # Window fetched by the should-run gate (_collect_scoped_interactions_for_precheck), # stashed so the billing path (_extraction_input_text) can reuse it instead of # re-querying storage. None when the gate did not run (bypass paths). @@ -624,6 +626,7 @@ def compute_generation(self, request: TRequest) -> GenerationComputePlan | None: self._last_extraction_run_ids = [] self._last_token_totals = None self._last_bookmark_advance = None + self._last_model_provenance = None result = self._execute_extractor(prepared.extractor_config, prepared.identifier) generated_count = self._count_generated_results(result) diff --git a/reflexio/server/services/deferred_learning_plan.py b/reflexio/server/services/deferred_learning_plan.py index 9c91ee289..eb0a87cd8 100644 --- a/reflexio/server/services/deferred_learning_plan.py +++ b/reflexio/server/services/deferred_learning_plan.py @@ -13,10 +13,12 @@ if TYPE_CHECKING: from reflexio.models.api_schema.domain.entities import ( + LineageContext, UserPlaybook, UserProfile, ) from reflexio.models.api_schema.service_schemas import Interaction + from reflexio.server.llm._litellm_types import ModelProvenance from reflexio.server.llm.token_accounting import RunTokenTotals from reflexio.server.services.base_generation_service import ( BaseGenerationService, @@ -70,6 +72,7 @@ class ProfileWritePlan: request_id: str new_profiles: list[UserProfile] superseded_ids: list[str] + lineage_contexts: list[LineageContext] = field(default_factory=list) @dataclass @@ -115,6 +118,8 @@ class PlaybookWritePlan: new_playbooks: list[UserPlaybook] superseded_ids: list[int] merge_groups: list[tuple[int, list[int]]] + lineage_contexts: list[LineageContext] = field(default_factory=list) + consolidation_provenance: ModelProvenance | None = None @dataclass diff --git a/reflexio/server/services/extraction/outcome.py b/reflexio/server/services/extraction/outcome.py index 8529f9941..60f991e88 100644 --- a/reflexio/server/services/extraction/outcome.py +++ b/reflexio/server/services/extraction/outcome.py @@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Literal if TYPE_CHECKING: + from reflexio.server.llm._litellm_types import ModelProvenance from reflexio.server.llm.token_accounting import RunTokenTotals from reflexio.server.services.deferred_learning_plan import ( ExtractorBookmarkAdvance, @@ -23,6 +24,7 @@ class ExtractionOutcome[T]: # The stride-bookmark advance the extractor no longer applies itself (F1); # applied downstream in persist (durable) or ``.run()``'s persist half. bookmark_advance: ExtractorBookmarkAdvance | None = None + model_provenance: ModelProvenance | None = None @classmethod def completed( @@ -32,6 +34,7 @@ def completed( run_id: str | None = None, token_totals: RunTokenTotals | None = None, bookmark_advance: ExtractorBookmarkAdvance | None = None, + model_provenance: ModelProvenance | None = None, ) -> ExtractionOutcome[T]: return cls( status="completed", @@ -39,6 +42,7 @@ def completed( run_id=run_id, token_totals=token_totals, bookmark_advance=bookmark_advance, + model_provenance=model_provenance, ) @classmethod diff --git a/reflexio/server/services/extraction/resumable_agent.py b/reflexio/server/services/extraction/resumable_agent.py index db6cfcb64..e59d507df 100644 --- a/reflexio/server/services/extraction/resumable_agent.py +++ b/reflexio/server/services/extraction/resumable_agent.py @@ -3,13 +3,14 @@ from __future__ import annotations import logging -from dataclasses import dataclass, replace +from dataclasses import asdict, dataclass, replace from datetime import UTC, datetime from typing import TYPE_CHECKING, Any from pydantic import BaseModel from reflexio.models.api_schema.internal_schema import RequestInteractionDataModel +from reflexio.server.llm._litellm_types import ModelProvenance from reflexio.server.llm.litellm_client import LiteLLMClient from reflexio.server.llm.model_defaults import ModelRole from reflexio.server.llm.tools import Tool, ToolLoopTrace, ToolRegistry, run_tool_loop @@ -76,6 +77,47 @@ class AgentRunResult: messages: list[dict[str, Any]] trace: ToolLoopTrace finished_reason: str + model_provenance: ModelProvenance | None = None + + +def encode_committed_output( + output: BaseModel, provenance: ModelProvenance | None +) -> dict[str, Any]: + """Persist output with provenance while old raw payloads remain readable.""" + return { + "_reflexio_envelope_version": 1, + "output": output.model_dump(), + "model_provenance": asdict(provenance) if provenance is not None else None, + } + + +def decode_committed_output( + committed_output: dict[str, Any], +) -> tuple[dict[str, Any], ModelProvenance | None]: + """Read the new envelope or a pre-provenance raw structured payload.""" + if "_reflexio_envelope_version" not in committed_output: + return committed_output, None + version = committed_output["_reflexio_envelope_version"] + if version != 1: + raise ValueError(f"Unsupported committed output envelope version: {version!r}") + output = committed_output.get("output") + if not isinstance(output, dict): + raise ValueError( + "Corrupt v1 committed output envelope: output must be an object" + ) + raw_provenance = committed_output.get("model_provenance") + if raw_provenance is None: + return output, None + if not isinstance(raw_provenance, dict): + raise ValueError( + "Corrupt v1 committed output envelope: model_provenance must be an object or null" + ) + try: + return output, ModelProvenance(**raw_provenance) + except (TypeError, ValueError) as exc: + raise ValueError( + "Corrupt v1 committed output envelope: invalid model_provenance" + ) from exc def _format_resolved_tool_result(record: PendingToolCallRecord) -> str: @@ -354,7 +396,11 @@ def _run( ) output = result.structured_output - committed_output = output.model_dump() if output is not None else None + committed_output = ( + encode_committed_output(output, result.provenance) + if output is not None + else None + ) active_statuses = (AgentRunStatus.RUNNING, AgentRunStatus.RESUMING) if ( result.finished_reason == "structured_output" @@ -390,6 +436,7 @@ def _run( messages=result.messages, trace=result.trace, finished_reason="late_output_discarded", + model_provenance=None, ) logger.info( "event=extraction_agent_finished org_id=%s user_id=%s " @@ -447,4 +494,5 @@ def _run( messages=result.messages, trace=result.trace, finished_reason=result.finished_reason, + model_provenance=result.provenance, ) diff --git a/reflexio/server/services/extraction/resume_worker.py b/reflexio/server/services/extraction/resume_worker.py index 0645cc896..3e1eaef06 100644 --- a/reflexio/server/services/extraction/resume_worker.py +++ b/reflexio/server/services/extraction/resume_worker.py @@ -14,6 +14,7 @@ from reflexio.models.api_schema.service_schemas import Interaction, Request from reflexio.models.config_schema import PlaybookConfig, ProfileExtractorConfig from reflexio.server.api_endpoints.request_context import RequestContext +from reflexio.server.llm._litellm_types import ModelProvenance from reflexio.server.llm.litellm_client import LiteLLMClient, LiteLLMConfig from reflexio.server.llm.model_defaults import ModelRole, resolve_model_name from reflexio.server.services.extraction.agent_run_records import build_scope_hash @@ -27,6 +28,7 @@ AgentRunResult, ResumableExtractionAgent, create_pending_info_tools_for_extractor_kind, + decode_committed_output, ) from reflexio.server.services.playbook.components.extractor import PlaybookExtractor from reflexio.server.services.playbook.playbook_service_utils import ( @@ -245,7 +247,9 @@ def run_once(self) -> AgentRunRecord | None: raise ResumeWorkerError( f"Run {run.id} has no resolved, unconsumed tool calls" ) - items, pending_tool_call_ids = self._resume_run(run, resolved_calls) + items, pending_tool_call_ids, model_provenance = self._resume_run( + run, resolved_calls + ) except Exception as exc: with sentry_tags( subsystem="extraction", @@ -272,7 +276,7 @@ def run_once(self) -> AgentRunRecord | None: try: self.storage.update_agent_run_status(run.id, AgentRunStatus.FINALIZING) - self._finalize_items(run, items) + self._finalize_items(run, items, model_provenance=model_provenance) self._schedule_finalized_tagging(run) self.storage.consume_run_tool_dependencies(run.id) finalized_status = ( @@ -315,8 +319,10 @@ def _retry_finalization(self, run: AgentRunRecord) -> AgentRunRecord | None: config = self.request_context.configurator.get_config() pending_config = config.pending_tool_call_config try: - items, pending_tool_call_ids = self._items_from_committed_output(run) - self._finalize_items(run, items) + items, pending_tool_call_ids, model_provenance = ( + self._items_from_committed_output(run) + ) + self._finalize_items(run, items, model_provenance=model_provenance) self._schedule_finalized_tagging(run) self.storage.consume_run_tool_dependencies(run.id) finalized_status = ( @@ -392,7 +398,7 @@ def _resume_run( self, run: AgentRunRecord, resolved_calls: list[PendingToolCallRecord], - ) -> tuple[list[Any], list[str]]: + ) -> tuple[list[Any], list[str], ModelProvenance | None]: request_interaction_data_models = _request_interaction_models_from_ids( self.storage, run.binding.source_interaction_ids, @@ -425,7 +431,7 @@ def _resume_profile( extractor_config: ProfileExtractorConfig | PlaybookConfig, request_interaction_data_models: list[RequestInteractionDataModel], resolved_calls: list[PendingToolCallRecord], - ) -> tuple[list[Any], list[str]]: + ) -> tuple[list[Any], list[str], ModelProvenance | None]: if not isinstance(extractor_config, ProfileExtractorConfig): raise ResumeWorkerError("Expected profile extractor config") if run.binding.user_id is None: @@ -498,6 +504,7 @@ def _resume_profile( source_interaction_ids=source_interaction_ids, ), result.pending_tool_call_ids, + result.model_provenance, ) def _resume_playbook( @@ -506,7 +513,7 @@ def _resume_playbook( extractor_config: ProfileExtractorConfig | PlaybookConfig, request_interaction_data_models: list[RequestInteractionDataModel], resolved_calls: list[PendingToolCallRecord], - ) -> tuple[list[Any], list[str]]: + ) -> tuple[list[Any], list[str], ModelProvenance | None]: if not isinstance(extractor_config, PlaybookConfig): raise ResumeWorkerError("Expected playbook extractor config") @@ -587,6 +594,7 @@ def _resume_playbook( source_interaction_ids=source_interaction_ids, ), result.pending_tool_call_ids, + result.model_provenance, ) def _messages_with_prior_knowledge( @@ -643,7 +651,7 @@ def _resume_agent( def _items_from_committed_output( self, run: AgentRunRecord, - ) -> tuple[list[Any], list[str]]: + ) -> tuple[list[Any], list[str], ModelProvenance | None]: if run.committed_output is None: raise ResumeWorkerError( f"Run {run.id} cannot retry finalization without committed output" @@ -655,19 +663,22 @@ def _items_from_committed_output( fallback_agent_version=run.binding.agent_version, ) extractor_config = _select_current_extractor_config(self.request_context, run) + output, model_provenance = decode_committed_output(run.committed_output) if run.binding.extractor_kind == "profile": - return self._profile_items_from_output( + items, pending_ids = self._profile_items_from_output( run, extractor_config, - run.committed_output, + output, ) + return items, pending_ids, model_provenance if run.binding.extractor_kind == "playbook": - return self._playbook_items_from_output( + items, pending_ids = self._playbook_items_from_output( run, extractor_config, request_interaction_data_models, - run.committed_output, + output, ) + return items, pending_ids, model_provenance raise ResumeWorkerError( f"Unsupported extractor kind {run.binding.extractor_kind!r}" ) @@ -750,7 +761,13 @@ def _playbook_items_from_output( run.pending_tool_call_ids, ) - def _finalize_items(self, run: AgentRunRecord, items: list[Any]) -> None: + def _finalize_items( + self, + run: AgentRunRecord, + items: list[Any], + *, + model_provenance: ModelProvenance | None = None, + ) -> None: if run.binding.extractor_kind == "profile": service = ProfileGenerationService( llm_client=self.client, @@ -763,7 +780,7 @@ def _finalize_items(self, run: AgentRunRecord, items: list[Any]) -> None: auto_run=False, force_extraction=True, ) - service._finalize_extracted_items(items) + service._finalize_extracted_items(items, model_provenance=model_provenance) self._record_finalized_learnings(run, items, entity_type="profile") return if run.binding.extractor_kind == "playbook": @@ -779,7 +796,7 @@ def _finalize_items(self, run: AgentRunRecord, items: list[Any]) -> None: auto_run=False, force_extraction=True, ) - service._finalize_extracted_items(items) + service._finalize_extracted_items(items, model_provenance=model_provenance) self._record_finalized_learnings(run, items, entity_type="user_playbook") return raise ResumeWorkerError( diff --git a/reflexio/server/services/playbook/components/aggregator.py b/reflexio/server/services/playbook/components/aggregator.py index 1cfda5b57..8d7603f76 100644 --- a/reflexio/server/services/playbook/components/aggregator.py +++ b/reflexio/server/services/playbook/components/aggregator.py @@ -10,6 +10,7 @@ if TYPE_CHECKING: import numpy as np +from reflexio.models.api_schema.domain.entities import LineageContext from reflexio.models.api_schema.service_schemas import ( AgentPlaybook, AgentPlaybookSourceWindow, @@ -21,6 +22,7 @@ PlaybookAggregatorConfig, ) from reflexio.server.api_endpoints.request_context import RequestContext +from reflexio.server.llm._litellm_types import ModelProvenance from reflexio.server.llm.litellm_client import LiteLLMClient from reflexio.server.services.operation_state_utils import OperationStateManager from reflexio.server.services.playbook.aggregation_prompt_processing import ( @@ -47,6 +49,7 @@ ensure_playbook_content, ) from reflexio.server.services.service_utils import log_model_response +from reflexio.server.services.storage.storage_base import AGGREGATE_REASON_PREFIX from reflexio.server.tracing import capture_anomaly, sentry_tags from reflexio.server.usage_metrics import record_usage_event @@ -476,7 +479,7 @@ def run(self, playbook_aggregator_request: PlaybookAggregatorRequest) -> dict: existing_playbooks, direction_overlap_threshold=playbook_aggregator_config.direction_overlap_threshold, ) - new_playbooks = [playbook for playbook, _ in generated_pairs] + new_playbooks = [playbook for playbook, _, _ in generated_pairs] previous_fingerprints_for_changed_clusters = {} changed_fps_by_previous_fp = {} @@ -558,19 +561,27 @@ def run(self, playbook_aggregator_request: PlaybookAggregatorRequest) -> dict: # Save each playbook + its aggregate event atomically, then assign # fingerprints and source-windows for the saved row. - for playbook, cluster_playbooks in generated_pairs: + for playbook, cluster_playbooks, provenance in generated_pairs: run_mode = "full_archive" if full_archive else "incremental" member_ids = [ str(fb.user_playbook_id) for fb in cluster_playbooks if fb.user_playbook_id ] - saved_fb = self.storage.save_agent_playbook_with_aggregate_event( # type: ignore[reportOptionalMemberAccess] - playbook, - source_ids=member_ids, - request_id=_run_id, - run_mode=run_mode, - ) + saved_fb = self.storage.save_agent_playbooks( # type: ignore[reportOptionalMemberAccess] + [playbook], + lineage_contexts=[ + LineageContext( + op_kind="aggregate", + actor="aggregator", + request_id=_run_id, + source_ids=member_ids, + reason=f"{AGGREGATE_REASON_PREFIX}{run_mode}", + model_name=provenance.model_name if provenance else None, + provider=provenance.provider if provenance else None, + ) + ], + )[0] saved_playbook_list.append(saved_fb) if saved_fb and saved_fb.agent_playbook_id: fp_key = self._compute_cluster_fingerprint(cluster_playbooks) @@ -784,7 +795,7 @@ def _record_learnings_generated( Prefers one event per learning id (entity-backed) when every saved playbook in this run carries a durable ``agent_playbook_id`` — the - common case, since ``save_agent_playbook_with_aggregate_event`` + common case, since ``save_agent_playbooks`` raises rather than returning a partial row. Falls back to the count-based aggregate event when ``learning_ids`` is short of ``total_count`` (a falsy/unset id slipped through), mirroring @@ -939,9 +950,11 @@ def _generate_playbooks_with_source_clusters( clusters: dict[int, list[UserPlaybook]], existing_approved_playbooks: list[AgentPlaybook], direction_overlap_threshold: float = 0.6, - ) -> list[tuple[AgentPlaybook, list[UserPlaybook]]]: - """Generate agent playbooks while preserving their exact source cluster.""" - new_playbooks: list[tuple[AgentPlaybook, list[UserPlaybook]]] = [] + ) -> list[tuple[AgentPlaybook, list[UserPlaybook], ModelProvenance | None]]: + """Generate playbooks with their exact source cluster and provenance.""" + new_playbooks: list[ + tuple[AgentPlaybook, list[UserPlaybook], ModelProvenance | None] + ] = [] approved_playbooks_str = ( "\n".join([f"- {fb.content}" for fb in existing_approved_playbooks]) if existing_approved_playbooks @@ -965,14 +978,15 @@ def _generate_playbooks_with_source_clusters( for playbook in cluster_playbooks ] - playbook = self._generate_playbook_from_cluster( + generated = self._generate_playbook_from_cluster( prompt_cluster_playbooks, approved_playbooks_str, direction_overlap_threshold=direction_overlap_threshold, processing_context=processing_context, ) - if playbook is not None: - new_playbooks.append((playbook, cluster_playbooks)) + if generated is not None: + playbook, provenance = generated + new_playbooks.append((playbook, cluster_playbooks, provenance)) return new_playbooks def _enqueue_playbook_optimization( @@ -1020,7 +1034,7 @@ def _generate_playbook_from_cluster( existing_approved_playbooks_str: str, direction_overlap_threshold: float = 0.6, processing_context: AggregationPromptProcessingContext | None = None, - ) -> AgentPlaybook | None: + ) -> tuple[AgentPlaybook, ModelProvenance | None] | None: """ Generate a playbook from a cluster using structured JSON output. @@ -1030,7 +1044,7 @@ def _generate_playbook_from_cluster( direction_overlap_threshold: Token overlap threshold for grouping by direction Returns: - AgentPlaybook | None: Generated playbook, or None if no new playbook needed + Generated playbook and its provenance, or None if no new playbook is needed """ if not cluster_playbooks: return None @@ -1064,7 +1078,10 @@ def _generate_playbook_from_cluster( playbook = self._process_aggregation_response(response, cluster_playbooks) if playbook is None: return None - return playbook.model_copy(update={"playbook_metadata": "mock_generated"}) + return ( + playbook.model_copy(update={"playbook_metadata": "mock_generated"}), + None, + ) # Format raw playbooks for prompt using structured format raw_playbooks_str = self._format_structured_cluster_input( @@ -1089,12 +1106,14 @@ def _generate_playbook_from_cluster( ] try: - response = self.client.generate_chat_response( + completion = self.client.generate_chat_response_with_provenance( messages=messages, model=self.client.config.model, response_format=PlaybookAggregationOutput, parse_structured_output=True, ) + response = completion.value + model_provenance = completion.provenance if isinstance(response, PlaybookAggregationOutput): response, artifact_count = ( self._postproc._postprocess_aggregation_response( @@ -1120,7 +1139,10 @@ def _generate_playbook_from_cluster( ) return None - return self._process_aggregation_response(response, cluster_playbooks) + playbook = self._process_aggregation_response(response, cluster_playbooks) + if playbook is None: + return None + return playbook, model_provenance except Exception as exc: processed_error, artifact_count = ( self._postproc._postprocess_aggregation_output( diff --git a/reflexio/server/services/playbook/components/consolidator.py b/reflexio/server/services/playbook/components/consolidator.py index 27c0d81d8..8944a2e50 100644 --- a/reflexio/server/services/playbook/components/consolidator.py +++ b/reflexio/server/services/playbook/components/consolidator.py @@ -18,10 +18,10 @@ ) from reflexio.models.structured_output import StrictStructuredOutput from reflexio.server.api_endpoints.request_context import RequestContext +from reflexio.server.llm._litellm_types import ModelProvenance from reflexio.server.llm.litellm_client import ( LiteLLMClient, LiteLLMClientError, - StructuredOutputRepairError, ) from reflexio.server.services.deduplication_utils import ( BaseDeduplicator, @@ -470,6 +470,8 @@ def __init__( """ super().__init__(request_context, llm_client) self._dedup_config = dedup_config or DeduplicationConfig() + self.model_provenance: ModelProvenance | None = None + self.consolidated_output_indices: set[int] = set() def _get_prompt_id(self) -> str: """Get the prompt ID for playbook consolidation.""" @@ -738,17 +740,21 @@ def _validate_output(output: BaseModel) -> list[str]: return [] try: - response = self.client.generate_chat_response( + completion = self.client.generate_chat_response_with_provenance( messages=[{"role": "user", "content": prompt}], model=self.model_name, response_format=output_schema_class, structured_output_validator=_validate_output, ) - except (StructuredOutputRepairError, LiteLLMClientError): + except LiteLLMClientError as exc: if first_parsed_output is not None: + self.model_provenance = exc.first_parsed_provenance return first_parsed_output raise + self.model_provenance = completion.provenance + response = completion.value + log_model_response(logger, "Consolidation response", response) if not isinstance(response, PlaybookConsolidationOutput): @@ -795,6 +801,8 @@ def deduplicate( generation_request_id = _normalize_generation_request_id( generation_request_id, request_id=request_id ) + self.model_provenance = None + self.consolidated_output_indices = set() if agent_version is None: raise TypeError("agent_version is required") @@ -1054,6 +1062,10 @@ def _build_deduplicated_results( # in the final list is the current length of ``new_rows``. if merge_source_ids: merge_groups.append((len(new_rows), merge_source_ids)) + if not isinstance(decision, IndependentDecision): + self.consolidated_output_indices.update( + range(len(new_rows), len(new_rows) + len(rows)) + ) new_rows.extend(rows) handled_new_ids.update(marked_new_ids) self._bump_counter(result_counters, decision.kind) diff --git a/reflexio/server/services/playbook/components/extractor.py b/reflexio/server/services/playbook/components/extractor.py index 9827e0326..0675229b7 100644 --- a/reflexio/server/services/playbook/components/extractor.py +++ b/reflexio/server/services/playbook/components/extractor.py @@ -82,6 +82,7 @@ def __init__( self.agent_context: str = agent_context self._last_resumable_run_id: str | None = None self._last_resumable_token_totals: RunTokenTotals | None = None + self._last_model_provenance = None # Get LLM config overrides from configuration config = self.request_context.configurator.get_config() @@ -224,6 +225,7 @@ def run(self) -> list[UserPlaybook] | ExtractionOutcome[UserPlaybook]: run_id=self._last_resumable_run_id, token_totals=self._last_resumable_token_totals, bookmark_advance=bookmark_advance, + model_provenance=self._last_model_provenance, ) def extract_playbook_entries( @@ -319,6 +321,7 @@ def extract_playbook_entries( ) self._last_resumable_run_id = result.run_id self._last_resumable_token_totals = sum_trace_tokens(result.trace) + self._last_model_provenance = result.model_provenance if not isinstance(result.output, StructuredPlaybookList): logger.warning( "Playbook extraction did not finish: %s", diff --git a/reflexio/server/services/playbook/playbook_edit_apply.py b/reflexio/server/services/playbook/playbook_edit_apply.py index 44e301d06..7dbc784b0 100644 --- a/reflexio/server/services/playbook/playbook_edit_apply.py +++ b/reflexio/server/services/playbook/playbook_edit_apply.py @@ -12,6 +12,10 @@ from reflexio.server.services.storage.storage_base import BaseStorage +class _LostSupersedeRaceError(Exception): + """Internal signal used to roll back a provisional successor.""" + + def apply_playbook_edit( storage: "BaseStorage", *, @@ -20,6 +24,7 @@ def apply_playbook_edit( source: str, request_id: str, skip_embedding: bool = False, + revise_context: LineageContext | None = None, ) -> int: """Insert a replacement playbook then atomically supersede the incumbent. @@ -29,12 +34,12 @@ def apply_playbook_edit( - Insert the new playbook as CURRENT. - Call ``supersede_record(incumbent_id → new_id)``, which only succeeds when the incumbent is still CURRENT (``status IS NULL``). - - If ``supersede_record`` returns ``False`` (incumbent already gone), delete - the just-inserted successor and return ``-1``. + - If ``supersede_record`` returns ``False`` (incumbent already gone), roll + back the transaction and return ``-1``. Args: storage: A BaseStorage instance providing ``save_user_playbooks``, - ``supersede_record``, and ``delete_user_playbooks_by_ids``. + ``supersede_record``, and ``commit_scope``. incumbent_id: ``user_playbook_id`` of the playbook being replaced. new_playbook: The replacement playbook (inserted as CURRENT, i.e. ``status=None``). @@ -45,8 +50,7 @@ def apply_playbook_edit( immediately (before any storage write) when empty, preventing orphaned successor rows. skip_embedding: Forwarded to ``save_user_playbooks``. Defaults to - ``False`` (recompute the embedding at write time — what every online - / offline-tuner caller relies on). + ``False`` (precompute the embedding before opening the transaction). Returns: The ``user_playbook_id`` of the newly inserted playbook, or ``-1`` if @@ -60,19 +64,24 @@ def apply_playbook_edit( "apply_playbook_edit: request_id must be non-empty (operation-run correlation id)" ) new_playbook.source = source - storage.save_user_playbooks([new_playbook], skip_embedding=skip_embedding) - new_id: int = new_playbook.user_playbook_id + if not skip_embedding: + storage.precompute_user_playbook_embeddings([new_playbook]) - ctx = LineageContext(op_kind="revise", actor=source, request_id=request_id) - superseded = storage.supersede_record( - entity_type="user_playbook", - incumbent_id=str(incumbent_id), - successor_id=str(new_id), - context=ctx, + ctx = revise_context or LineageContext( + op_kind="revise", actor=source, request_id=request_id ) - if not superseded: - # lost the race: delete the just-inserted successor so no orphan CURRENT row - # remains. It was never live, so this is a rollback — not an audited erasure. - storage.delete_user_playbooks_by_ids([new_id], emit_hard_delete=False) + try: + with storage.commit_scope(): + storage.save_user_playbooks([new_playbook], skip_embedding=True) + new_id = new_playbook.user_playbook_id + if not storage.supersede_record( + entity_type="user_playbook", + incumbent_id=str(incumbent_id), + successor_id=str(new_id), + context=ctx, + ): + raise _LostSupersedeRaceError + except _LostSupersedeRaceError: + new_playbook.user_playbook_id = 0 return -1 return new_id diff --git a/reflexio/server/services/playbook/service.py b/reflexio/server/services/playbook/service.py index 4a9818bde..3a6f24c79 100644 --- a/reflexio/server/services/playbook/service.py +++ b/reflexio/server/services/playbook/service.py @@ -11,6 +11,7 @@ from reflexio.server.services.deferred_learning_plan import GenerationComputePlan from reflexio.server.services.storage.storage_base import BaseStorage +from reflexio.models.api_schema.domain.entities import LineageContext from reflexio.models.api_schema.internal_schema import RequestInteractionDataModel from reflexio.models.api_schema.service_schemas import ( DowngradeUserPlaybooksResponse, @@ -23,6 +24,7 @@ UserPlaybook, ) from reflexio.models.config_schema import PlaybookConfig +from reflexio.server.llm._litellm_types import ModelProvenance from reflexio.server.services.base_generation_service import ( BaseGenerationService, StatusChangeOperation, @@ -338,6 +340,8 @@ def _resolve_write_plan( self.service_config.agent_version, # type: ignore[reportOptionalMemberAccess] user_id=self.service_config.user_id, # type: ignore[reportOptionalMemberAccess] ) + consolidation_provenance = consolidator.model_provenance + consolidated_output_indices = consolidator.consolidated_output_indices logger.info( "User playbook entries after deduplication: %d", len(deduplicated_playbooks), @@ -366,6 +370,27 @@ def _resolve_write_plan( # persist half passes skip_embedding=True so no embedding runs in the fence. self.storage.precompute_user_playbook_embeddings(all_playbooks) # type: ignore[reportOptionalMemberAccess] + lineage_contexts: list[LineageContext] = [] + for index, _playbook in enumerate(all_playbooks): + provenance = ( + consolidation_provenance + if index in consolidated_output_indices + else self._last_model_provenance + ) + lineage_contexts.append( + LineageContext( + op_kind="create", + actor=( + "consolidator" + if index in consolidated_output_indices + else "extractor" + ), + request_id=generation_request_id, + model_name=provenance.model_name if provenance else None, + provider=provenance.provider if provenance else None, + ) + ) + return PlaybookWritePlan( request_id=generation_request_id, output_pending_status=self.output_pending_status, @@ -373,6 +398,8 @@ def _resolve_write_plan( new_playbooks=all_playbooks, superseded_ids=existing_ids_to_delete, merge_groups=merge_groups, + lineage_contexts=lineage_contexts, + consolidation_provenance=consolidation_provenance, ) def _persist_write_plan(self, plan: PlaybookWritePlan) -> None: @@ -391,13 +418,16 @@ def _persist_write_plan(self, plan: PlaybookWritePlan) -> None: return try: self.storage.save_user_playbooks( # type: ignore[reportOptionalMemberAccess] - plan.new_playbooks, skip_embedding=True + plan.new_playbooks, + skip_embedding=True, + lineage_contexts=plan.lineage_contexts, ) self._apply_consolidation_lineage( plan.new_playbooks, plan.merge_groups, plan.superseded_ids, request_id=plan.request_id, + model_provenance=plan.consolidation_provenance, ) except Exception as e: logger.error( @@ -438,7 +468,12 @@ def emit_generation_side_effects(self, plan: GenerationComputePlan) -> None: if write_plan is not None: self._dispatch_playbook_schedulers(write_plan) - def _finalize_extracted_items(self, all_playbooks: list[UserPlaybook]) -> None: + def _finalize_extracted_items( + self, + all_playbooks: list[UserPlaybook], + *, + model_provenance: ModelProvenance | None = None, + ) -> None: """Permanent V3 wrapper: compute→persist→schedulers together (no fence). Kept for the synchronous resume/manual callers @@ -448,6 +483,8 @@ def _finalize_extracted_items(self, all_playbooks: list[UserPlaybook]) -> None: ``commit_scope`` — then dispatches the same off-thread schedulers, so the result is identical to the pre-split monolith. """ + if model_provenance is not None: + self._last_model_provenance = model_provenance plan = self._resolve_write_plan([all_playbooks]) if plan is None: return @@ -461,6 +498,7 @@ def _apply_consolidation_lineage( existing_ids_to_delete: list[int], *, request_id: str, + model_provenance: ModelProvenance | None = None, ) -> None: """Materialize consolidation merges as lineage tombstones. @@ -482,8 +520,6 @@ def _apply_consolidation_lineage( rather than read off ``self.service_config`` so persist stays decoupled from the mutable service config on the fenced path. """ - from reflexio.models.api_schema.domain.entities import LineageContext - generation_request_id = request_id merged_source_ids: set[int] = set() for survivor_idx, source_ids in merge_groups: @@ -499,6 +535,10 @@ def _apply_consolidation_lineage( source_ids=[str(s) for s in source_ids], reason="dedup-merge", request_id=generation_request_id, + model_name=( + model_provenance.model_name if model_provenance else None + ), + provider=model_provenance.provider if model_provenance else None, ), ) diff --git a/reflexio/server/services/playbook_optimizer/optimizer.py b/reflexio/server/services/playbook_optimizer/optimizer.py index ffe17c8ce..4c3425262 100644 --- a/reflexio/server/services/playbook_optimizer/optimizer.py +++ b/reflexio/server/services/playbook_optimizer/optimizer.py @@ -23,6 +23,7 @@ from reflexio.server.services.playbook.aggregation_trigger import ( maybe_trigger_user_playbook_aggregation, ) +from reflexio.server.services.playbook.playbook_edit_apply import apply_playbook_edit from reflexio.server.tracing import sentry_tags from .assistant_webhook import AssistantCallable, LocalScriptAssistant, WebhookAssistant @@ -559,8 +560,7 @@ def _supersede_user_playbook( incumbent is no longer CURRENT (lost race / already superseded). Args: - storage: A storage instance implementing ``save_user_playbooks``, - ``supersede_record``, and ``delete_user_playbooks_by_ids``. + storage: A storage instance implementing the canonical atomic edit path. incumbent: The current user playbook to replace. best_content: Content for the successor playbook. source: Provenance label written to the lineage event actor field. @@ -581,25 +581,24 @@ def _supersede_user_playbook( successor = incumbent.model_copy( update={"user_playbook_id": 0, "content": best_content, "status": None} ) - storage.save_user_playbooks([successor]) ctx = LineageContext( op_kind="revise", actor=source, request_id=request_id, ) - ok = storage.supersede_record( - entity_type="user_playbook", - incumbent_id=str(incumbent.user_playbook_id), - successor_id=str(successor.user_playbook_id), - context=ctx, + successor_id = apply_playbook_edit( + storage, + incumbent_id=incumbent.user_playbook_id, + new_playbook=successor, + source=source, + request_id=request_id, + revise_context=ctx, ) - if not ok: - # Lost CAS: remove the never-live successor without auditing it as erasure. - storage.delete_user_playbooks_by_ids( - [successor.user_playbook_id], emit_hard_delete=False - ) - return None - return successor.user_playbook_id + return None if successor_id == -1 else successor_id + + +class _LostAgentSupersedeRaceError(Exception): + """Internal rollback signal for an agent-playbook successor race.""" def _supersede_agent_playbook( @@ -618,7 +617,7 @@ def _supersede_agent_playbook( Args: storage: A storage instance implementing ``save_agent_playbooks``, - ``supersede_record``, and ``delete_agent_playbooks_by_ids``. + ``supersede_record``, and ``commit_scope``. incumbent: The current agent playbook to replace. best_content: Content for the successor playbook. source: Provenance label written to the lineage event actor field. @@ -646,24 +645,26 @@ def _supersede_agent_playbook( "playbook_metadata": playbook_metadata, } ) - saved = storage.save_agent_playbooks([successor]) - if not saved or not saved[0].agent_playbook_id: - return None - successor_id = saved[0].agent_playbook_id ctx = LineageContext( op_kind="revise", actor=source, request_id=request_id, ) - ok = storage.supersede_record( - entity_type="agent_playbook", - incumbent_id=str(incumbent.agent_playbook_id), - successor_id=str(successor_id), - context=ctx, - ) - if not ok: - # Lost CAS: remove the never-live successor without auditing it as erasure. - storage.delete_agent_playbooks_by_ids([successor_id], emit_hard_delete=False) + try: + with storage.commit_scope(): + saved = storage.save_agent_playbooks([successor]) + if not saved or not saved[0].agent_playbook_id: + raise _LostAgentSupersedeRaceError + successor_id = saved[0].agent_playbook_id + if not storage.supersede_record( + entity_type="agent_playbook", + incumbent_id=str(incumbent.agent_playbook_id), + successor_id=str(successor_id), + context=ctx, + ): + raise _LostAgentSupersedeRaceError + except _LostAgentSupersedeRaceError: + successor.agent_playbook_id = 0 return None return successor_id diff --git a/reflexio/server/services/profile/components/consolidator.py b/reflexio/server/services/profile/components/consolidator.py index 31f9d55ee..ebced17e6 100644 --- a/reflexio/server/services/profile/components/consolidator.py +++ b/reflexio/server/services/profile/components/consolidator.py @@ -15,6 +15,7 @@ from reflexio.models.api_schema.service_schemas import Status, UserProfile from reflexio.models.structured_output import StrictStructuredOutput from reflexio.server.api_endpoints.request_context import RequestContext +from reflexio.server.llm._litellm_types import ModelProvenance from reflexio.server.llm.litellm_client import ( LiteLLMClient, LiteLLMClientError, @@ -333,6 +334,9 @@ def __init__( """ super().__init__(request_context, llm_client) self.output_pending_status = output_pending_status + self.model_provenance: ModelProvenance | None = None + self.lineage_sources_by_profile_id: dict[str, list[str]] = {} + self.consolidated_output_indices: set[int] = set() def _get_prompt_id(self) -> str: """Get the prompt ID for profile deduplication.""" @@ -512,6 +516,10 @@ def deduplicate( Returns: Tuple of (deduplicated profiles, existing profile IDs to delete, superseded existing profiles) """ + self.model_provenance = None + self.lineage_sources_by_profile_id = {} + self.consolidated_output_indices = set() + # Check if mock mode is enabled if os.getenv("MOCK_LLM_RESPONSE", "").lower() == "true": logger.info("Mock mode: skipping deduplication") @@ -577,12 +585,14 @@ def _validate_output(output: BaseModel) -> list[str]: logger, "Profile deduplication", [{"role": "user", "content": prompt}] ) - response = self.client.generate_chat_response( + completion = self.client.generate_chat_response_with_provenance( messages=[{"role": "user", "content": prompt}], model=self.model_name, response_format=output_schema_class, structured_output_validator=_validate_output, ) + self.model_provenance = completion.provenance + response = completion.value log_model_response(logger, "Deduplication response", response) @@ -630,6 +640,9 @@ def _validate_output(output: BaseModel) -> list[str]: # drops out-of-range indices, skips a group it cannot resolve # without marking anything, and re-adds any unreferenced NEW # profile via the safety fallback. + # Ladder walk stamps first_parsed_provenance across all rungs so this + # matches first_parsed_output from the shared validator closure. + self.model_provenance = getattr(e, "first_parsed_provenance", None) logger.warning( "Falling back to the first parsed deduplication attempt after " "repair exhausted" @@ -825,7 +838,14 @@ def _build_deduplicated_results( status=template_profile.status, extractor_names=merged_extractor_names, ) + self.consolidated_output_indices.add(len(result_profiles)) result_profiles.append(merged_profile) + self.lineage_sources_by_profile_id[merged_profile.profile_id] = [ + str(existing_profiles[eidx].profile_id) + for eidx in group_existing_indices + if 0 <= eidx < len(existing_profiles) + and existing_profiles[eidx].profile_id + ] # Add unique NEW profiles for uid in dedup_output.unique_ids: diff --git a/reflexio/server/services/profile/components/extractor.py b/reflexio/server/services/profile/components/extractor.py index 588018aee..e44937ba9 100644 --- a/reflexio/server/services/profile/components/extractor.py +++ b/reflexio/server/services/profile/components/extractor.py @@ -99,6 +99,7 @@ def __init__( self.agent_context = agent_context self._last_resumable_run_id: str | None = None self._last_resumable_token_totals: RunTokenTotals | None = None + self._last_model_provenance = None # Get LLM config overrides from configuration config = self.request_context.configurator.get_config() @@ -272,6 +273,7 @@ def run(self) -> list[UserProfile] | ExtractionOutcome[UserProfile] | None: run_id=self._last_resumable_run_id, token_totals=self._last_resumable_token_totals, bookmark_advance=bookmark_advance, + model_provenance=self._last_model_provenance, ) def _convert_raw_to_user_profiles( @@ -409,6 +411,7 @@ def _generate_raw_updates_from_sessions( ) self._last_resumable_run_id = result.run_id self._last_resumable_token_totals = sum_trace_tokens(result.trace) + self._last_model_provenance = result.model_provenance if not isinstance(result.output, StructuredProfilesOutput): logger.warning( "Profile extraction did not finish: %s", result.finished_reason diff --git a/reflexio/server/services/profile/service.py b/reflexio/server/services/profile/service.py index fb7420022..4fafc70fd 100644 --- a/reflexio/server/services/profile/service.py +++ b/reflexio/server/services/profile/service.py @@ -11,6 +11,7 @@ from reflexio.server.api_endpoints.request_context import RequestContext from reflexio.server.llm.litellm_client import LiteLLMClient +from reflexio.models.api_schema.domain.entities import LineageContext from reflexio.models.api_schema.internal_schema import RequestInteractionDataModel from reflexio.models.api_schema.service_schemas import ( DowngradeProfilesResponse, @@ -23,6 +24,7 @@ UserProfile, ) from reflexio.models.config_schema import ProfileExtractorConfig +from reflexio.server.llm._litellm_types import ModelProvenance from reflexio.server.services.base_generation_service import ( BaseGenerationService, StatusChangeOperation, @@ -168,6 +170,9 @@ def _resolve_write_plan( all_new_profiles = [p for result in results if result for p in result] existing_ids_to_delete: list[str] = [] + consolidation_provenance = None + consolidation_sources: dict[str, list[str]] = {} + consolidated_output_indices: set[int] = set() # Always run deduplicator when there are new profiles if all_new_profiles: @@ -185,6 +190,9 @@ def _resolve_write_plan( all_new_profiles, user_id, generation_request_id ) ) + consolidation_provenance = consolidator.model_provenance + consolidation_sources = consolidator.lineage_sources_by_profile_id + consolidated_output_indices = consolidator.consolidated_output_indices logger.info( "Profile updates after deduplication: %d profiles, %d existing to delete", len(all_new_profiles), @@ -217,11 +225,34 @@ def _resolve_write_plan( if all_new_profiles: self.storage.precompute_profile_embeddings(all_new_profiles) # type: ignore[reportOptionalMemberAccess] + lineage_contexts: list[LineageContext] = [] + for index, profile in enumerate(all_new_profiles): + provenance = ( + consolidation_provenance + if index in consolidated_output_indices + else self._last_model_provenance + ) + lineage_contexts.append( + LineageContext( + op_kind="create", + actor=( + "consolidator" + if index in consolidated_output_indices + else "extractor" + ), + request_id=generation_request_id, + source_ids=consolidation_sources.get(profile.profile_id, []), + model_name=provenance.model_name if provenance else None, + provider=provenance.provider if provenance else None, + ) + ) + return ProfileWritePlan( user_id=user_id, request_id=generation_request_id, new_profiles=all_new_profiles, superseded_ids=existing_ids_to_delete, + lineage_contexts=lineage_contexts, ) def _persist_write_plan(self, plan: ProfileWritePlan) -> None: @@ -250,7 +281,10 @@ def _persist_write_plan(self, plan: ProfileWritePlan) -> None: if plan.new_profiles: try: self.storage.add_user_profile( # type: ignore[reportOptionalMemberAccess] - user_id, plan.new_profiles, skip_embedding=True + user_id, + plan.new_profiles, + skip_embedding=True, + lineage_contexts=plan.lineage_contexts, ) except Exception as e: with sentry_tags( @@ -298,7 +332,12 @@ def _persist_write_plan(self, plan: ProfileWritePlan) -> None: # _apply_consolidation_lineage raises here too. raise - def _finalize_extracted_items(self, all_new_profiles: list[UserProfile]) -> None: + def _finalize_extracted_items( + self, + all_new_profiles: list[UserProfile], + *, + model_provenance: ModelProvenance | None = None, + ) -> None: """Permanent V3 wrapper: compute-then-persist together (no external fence). Kept for the synchronous resume/manual callers @@ -307,6 +346,8 @@ def _finalize_extracted_items(self, all_new_profiles: list[UserProfile]) -> None (persist) split the durable worker uses — with no external ``commit_scope`` — so the result is identical to the pre-split monolith. """ + if model_provenance is not None: + self._last_model_provenance = model_provenance plan = self._resolve_write_plan([all_new_profiles]) if plan is not None: self._persist_write_plan(plan) diff --git a/reflexio/server/services/storage/sqlite_storage/_base.py b/reflexio/server/services/storage/sqlite_storage/_base.py index 448c8bef2..299852269 100644 --- a/reflexio/server/services/storage/sqlite_storage/_base.py +++ b/reflexio/server/services/storage/sqlite_storage/_base.py @@ -1220,6 +1220,8 @@ def _migrate_lineage_event_table(self) -> None: request_id TEXT NOT NULL DEFAULT '', reason TEXT NOT NULL DEFAULT '', created_at INTEGER NOT NULL, + model_name TEXT, + provider TEXT, UNIQUE (org_id, entity_type, entity_id, op, request_id) ); CREATE INDEX IF NOT EXISTS idx_lineage_entity @@ -1231,7 +1233,13 @@ def _migrate_lineage_event_table(self) -> None: "PRAGMA table_info(lineage_event)" ).fetchall() } - for col in ("from_status", "to_status", "status_namespace"): + for col in ( + "from_status", + "to_status", + "status_namespace", + "model_name", + "provider", + ): if col not in existing_cols: self.conn.execute( f"ALTER TABLE lineage_event ADD COLUMN {col} TEXT" # noqa: S608 @@ -2337,6 +2345,8 @@ def clear_user_data(self, user_id: str) -> dict[str, int]: from_status TEXT, to_status TEXT, status_namespace TEXT, + model_name TEXT, + provider TEXT, UNIQUE (org_id, entity_type, entity_id, op, request_id) ); CREATE INDEX IF NOT EXISTS idx_lineage_entity ON lineage_event (entity_type, entity_id); diff --git a/reflexio/server/services/storage/sqlite_storage/_lineage.py b/reflexio/server/services/storage/sqlite_storage/_lineage.py index 376bd5837..f7bdc80a6 100644 --- a/reflexio/server/services/storage/sqlite_storage/_lineage.py +++ b/reflexio/server/services/storage/sqlite_storage/_lineage.py @@ -61,6 +61,8 @@ def _append_event_stmt( from_status: str | None = None, to_status: str | None = None, status_namespace: str | None = None, + model_name: str | None = None, + provider: str | None = None, ) -> sqlite3.Cursor: """Insert a lineage event row; no-ops on (org_id, entity_type, entity_id, op, request_id) duplicate. @@ -70,8 +72,9 @@ def _append_event_stmt( "INSERT OR IGNORE INTO lineage_event " "(org_id, entity_type, entity_id, op, prov_relation, source_ids, " "actor, request_id, reason, created_at, " - "from_status, to_status, status_namespace) " - "VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?)", + "from_status, to_status, status_namespace, model_name, " + "provider) " + "VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)", ( org_id, entity_type, @@ -86,6 +89,8 @@ def _append_event_stmt( from_status, to_status, status_namespace, + model_name, + provider, ), ) @@ -158,6 +163,8 @@ def append_lineage_event(self, event: LineageEvent) -> int: from_status=event.from_status, to_status=event.to_status, status_namespace=event.status_namespace, + model_name=event.model_name, + provider=event.provider, ) if ( cur.rowcount == 0 @@ -234,6 +241,8 @@ def get_lineage_events( from_status=r["from_status"], to_status=r["to_status"], status_namespace=r["status_namespace"], + model_name=r["model_name"], + provider=r["provider"], ) for r in rows ] @@ -303,6 +312,8 @@ def merge_records( actor=context.actor, request_id=context.request_id, reason=context.reason, + model_name=context.model_name, + provider=context.provider, ) if self._own_transaction(): self.conn.commit() @@ -361,6 +372,8 @@ def supersede_record( actor=context.actor, request_id=context.request_id, reason=context.reason, + model_name=context.model_name, + provider=context.provider, ) if self._own_transaction(): self.conn.commit() diff --git a/reflexio/server/services/storage/sqlite_storage/playbook/_agent.py b/reflexio/server/services/storage/sqlite_storage/playbook/_agent.py index 9131d7497..7b6f7ef62 100644 --- a/reflexio/server/services/storage/sqlite_storage/playbook/_agent.py +++ b/reflexio/server/services/storage/sqlite_storage/playbook/_agent.py @@ -10,6 +10,7 @@ logger = logging.getLogger(__name__) from reflexio.models.api_schema.common import BlockingIssue +from reflexio.models.api_schema.domain.entities import LineageContext from reflexio.models.api_schema.retriever_schema import ( SearchAgentPlaybookRequest, ) @@ -23,9 +24,6 @@ from reflexio.server.services.storage.lifecycle_filters import ( validate_include_inactive, ) -from reflexio.server.services.storage.storage_base._playbook import ( - AGGREGATE_REASON_PREFIX, -) from .._base import ( _TOMBSTONE_STATUS_VALUES, @@ -93,6 +91,7 @@ class AgentPlaybookStoreMixin: _vec_upsert: Any _delete_playbook_search_rows: Any _own_transaction: Any + commit_scope: Any def _index_agent_playbook_fts_vec(self, ap: AgentPlaybook) -> None: """Update the FTS and vector indexes for a single agent playbook row. @@ -167,10 +166,34 @@ def _insert_agent_playbook_row( @SQLiteStorageBase.handle_exceptions def save_agent_playbooks( - self, agent_playbooks: list[AgentPlaybook] + self, + agent_playbooks: list[AgentPlaybook], + *, + lineage_contexts: list[LineageContext] | None = None, ) -> list[AgentPlaybook]: - saved: list[AgentPlaybook] = [] - for ap in agent_playbooks: + if lineage_contexts is not None and len(lineage_contexts) != len( + agent_playbooks + ): + raise ValueError("lineage_contexts must match agent_playbooks length") + if any( + context.op_kind not in {"create", "aggregate"} + for context in lineage_contexts or [] + ): + raise ValueError( + "agent playbook lineage context must use op_kind='create' or 'aggregate'" + ) + if any( + context.op_kind == "aggregate" + and not (context.request_id and context.request_id.strip()) + for context in lineage_contexts or [] + ): + raise ValueError("agent playbook aggregate lineage requires request_id") + + contexts = lineage_contexts or [ + LineageContext(op_kind="create") for _ap in agent_playbooks + ] + rows: list[tuple[AgentPlaybook, LineageContext, str]] = [] + for ap, context in zip(agent_playbooks, contexts, strict=True): embedding_text = ap.trigger or ap.content if self._should_expand_documents(): with ThreadPoolExecutor(max_workers=2) as executor: @@ -181,94 +204,39 @@ def save_agent_playbooks( else: ap.embedding = self._get_embedding(embedding_text) - created_at_iso = _epoch_to_iso(ap.created_at) - with self._lock: - self._insert_agent_playbook_row(self.conn, ap, created_at_iso) - if self._own_transaction(): - self.conn.commit() - - self._index_agent_playbook_fts_vec(ap) - saved.append(ap) - return saved + rows.append((ap, context, _epoch_to_iso(ap.created_at))) - @SQLiteStorageBase.handle_exceptions - def save_agent_playbook_with_aggregate_event( - self, - agent_playbook: AgentPlaybook, - *, - source_ids: list[str], - request_id: str, - run_mode: str, - ) -> AgentPlaybook: - """Persist an agent playbook AND its ``op=aggregate`` lineage event atomically. - - The INSERT and the event are committed in a single transaction — if either - fails, both roll back. The event is the sole record of the run->playbook - membership for reconstruction, so atomicity is critical. - - Args: - agent_playbook (AgentPlaybook): The playbook to persist. - source_ids (list[str]): IDs of the source entities that produced this playbook. - request_id (str): The aggregation run ID. - run_mode (str): Aggregation run mode (e.g. ``full_archive`` or ``incremental``). - - Returns: - AgentPlaybook: The saved playbook with ``agent_playbook_id`` populated. - - Raises: - ValueError: If ``request_id`` is empty (would produce an unreconstructable event). - """ - if not request_id or not request_id.strip(): - raise ValueError( - "save_agent_playbook_with_aggregate_event requires a non-empty request_id" - ) - ap = agent_playbook - embedding_text = ap.trigger or ap.content - if self._should_expand_documents(): - with ThreadPoolExecutor(max_workers=2) as executor: - emb_future = executor.submit(self._get_embedding, embedding_text) - exp_future = executor.submit(self._expand_document, embedding_text) - ap.embedding = emb_future.result(timeout=15) - ap.expanded_terms = exp_future.result(timeout=15) - else: - ap.embedding = self._get_embedding(embedding_text) + with self.commit_scope(): + for ap, context, created_at_iso in rows: + with self._lock: + self._insert_agent_playbook_row(self.conn, ap, created_at_iso) + is_aggregate = context.op_kind == "aggregate" + _append_event_stmt( + self.conn, + org_id=self.org_id, + entity_type="agent_playbook", + entity_id=str(ap.agent_playbook_id), + op=context.op_kind, + prov="wasDerivedFrom" if is_aggregate else "wasGeneratedBy", + source_ids=context.source_ids, + actor=context.actor, + request_id=context.request_id + or f"{context.op_kind}_{ap.agent_playbook_id}", + reason=context.reason, + model_name=context.model_name, + provider=context.provider, + ) - created_at_iso = _epoch_to_iso(ap.created_at) - with self._lock: - own_txn = self._own_transaction() + for ap, _context, _created_at_iso in rows: try: - self._insert_agent_playbook_row(self.conn, ap, created_at_iso) - _append_event_stmt( - self.conn, - org_id=self.org_id, - entity_type="agent_playbook", - entity_id=str(ap.agent_playbook_id), - op="aggregate", - prov="wasDerivedFrom", - source_ids=source_ids, - actor="aggregator", - request_id=request_id, - reason=f"{AGGREGATE_REASON_PREFIX}{run_mode}", - ) - if own_txn: - self.conn.commit() + self._index_agent_playbook_fts_vec(ap) except Exception: - if own_txn: - self.conn.rollback() - raise - - # FTS/vec indexing AFTER commit — these helpers self-commit and must - # not be interleaved inside the atomic transaction above. - # Index failure does NOT invalidate the committed row+event; the index - # is reconstructable from the authoritative row. - try: - self._index_agent_playbook_fts_vec(ap) - except Exception: - logger.exception( - "FTS/vec indexing failed for agent_playbook %s (row committed, index skipped)", - ap.agent_playbook_id, - ) - return ap + logger.exception( + "FTS/vec indexing failed for agent_playbook %s " + "(row committed, index skipped)", + ap.agent_playbook_id, + ) + return agent_playbooks @SQLiteStorageBase.handle_exceptions def get_agent_playbooks( diff --git a/reflexio/server/services/storage/sqlite_storage/playbook/_user.py b/reflexio/server/services/storage/sqlite_storage/playbook/_user.py index db12cee05..dd10e65f7 100644 --- a/reflexio/server/services/storage/sqlite_storage/playbook/_user.py +++ b/reflexio/server/services/storage/sqlite_storage/playbook/_user.py @@ -7,6 +7,7 @@ from typing import Any from reflexio.models.api_schema.common import BlockingIssue +from reflexio.models.api_schema.domain.entities import LineageContext from reflexio.models.api_schema.retriever_schema import SearchUserPlaybookRequest from reflexio.models.api_schema.service_schemas import Status, UserPlaybook from reflexio.models.config_schema import SearchMode, SearchOptions @@ -77,6 +78,7 @@ class UserPlaybookStoreMixin: _subject_ref_for_user_id: Any _assert_subject_writable_locked: Any _own_transaction: Any + commit_scope: Any def _subject_ref_from_user_playbook_row(self, row: sqlite3.Row) -> str: subject_ref = row["governance_subject_ref"] @@ -134,8 +136,27 @@ def save_user_playbooks( user_playbooks: list[UserPlaybook], *, skip_embedding: bool = False, + lineage_contexts: list[LineageContext] | None = None, ) -> None: - for up in user_playbooks: + if lineage_contexts is not None and len(lineage_contexts) != len( + user_playbooks + ): + raise ValueError("lineage_contexts must match user_playbooks length") + if any(context.op_kind != "create" for context in lineage_contexts or []): + raise ValueError( + "user playbook create lineage context must use op_kind='create'" + ) + contexts = lineage_contexts or [ + LineageContext( + op_kind="create", + actor=up.source or "", + source_ids=[str(value) for value in up.source_interaction_ids], + request_id=up.request_id, + ) + for up in user_playbooks + ] + rows: list[tuple[UserPlaybook, LineageContext, str, str]] = [] + for up, lineage_context in zip(user_playbooks, contexts, strict=True): subject_ref = self._subject_ref_for_user_id(up.user_id) with self._lock: self._assert_subject_writable_locked(subject_ref) @@ -145,13 +166,13 @@ def save_user_playbooks( # out (embedding already set by precompute_user_playbook_embeddings). if not skip_embedding: self.precompute_user_playbook_embeddings([up]) + rows.append( + (up, lineage_context, subject_ref, _epoch_to_iso(up.created_at)) + ) - created_at_iso = _epoch_to_iso(up.created_at) - with self._lock: - own_txn = self._own_transaction() - try: - if own_txn: - self.conn.execute("BEGIN IMMEDIATE") + with self.commit_scope(): + for up, lineage_context, subject_ref, created_at_iso in rows: + with self._lock: self._assert_subject_writable_locked(subject_ref) cur = self.conn.execute( """INSERT INTO user_playbooks @@ -190,13 +211,25 @@ def save_user_playbooks( ) upid = cur.lastrowid or 0 up.user_playbook_id = upid - if own_txn: - self.conn.commit() - except Exception: - if own_txn: - self.conn.rollback() - raise + _append_event_stmt( + self.conn, + org_id=self.org_id, + entity_type="user_playbook", + entity_id=str(upid), + op="create", + prov="wasGeneratedBy", + source_ids=lineage_context.source_ids, + actor=lineage_context.actor, + request_id=lineage_context.request_id + or up.request_id + or f"create_{upid}", + reason=lineage_context.reason, + model_name=lineage_context.model_name, + provider=lineage_context.provider, + ) + for up, _lineage_context, _subject_ref, _created_at_iso in rows: + upid = up.user_playbook_id fts_parts = [up.trigger or "", up.content or ""] if up.expanded_terms: fts_parts.append(up.expanded_terms) diff --git a/reflexio/server/services/storage/sqlite_storage/profiles/_profile_store.py b/reflexio/server/services/storage/sqlite_storage/profiles/_profile_store.py index 78ba9d8f2..4c3dbd2d1 100644 --- a/reflexio/server/services/storage/sqlite_storage/profiles/_profile_store.py +++ b/reflexio/server/services/storage/sqlite_storage/profiles/_profile_store.py @@ -11,6 +11,7 @@ from concurrent.futures import ThreadPoolExecutor from typing import Any +from reflexio.models.api_schema.domain.entities import LineageContext from reflexio.models.api_schema.service_schemas import ( DeleteUserProfileRequest, Status, @@ -81,6 +82,7 @@ class ProfileStoreMixin: _subject_ref_for_user_id: Any _assert_subject_writable_locked: Any _own_transaction: Any + commit_scope: Any def _subject_ref_from_profile_row(self, row: sqlite3.Row) -> str: subject_ref = row["governance_subject_ref"] @@ -236,8 +238,23 @@ def add_user_profile( user_profiles: list[UserProfile], *, skip_embedding: bool = False, + lineage_contexts: list[LineageContext] | None = None, ) -> None: - for profile in user_profiles: + if lineage_contexts is not None and len(lineage_contexts) != len(user_profiles): + raise ValueError("lineage_contexts must match user_profiles length") + if any(context.op_kind != "create" for context in lineage_contexts or []): + raise ValueError("profile create lineage context must use op_kind='create'") + contexts = lineage_contexts or [ + LineageContext( + op_kind="create", + actor=profile.source or "", + source_ids=[str(value) for value in profile.source_interaction_ids], + request_id=profile.generated_from_request_id, + ) + for profile in user_profiles + ] + rows: list[tuple[UserProfile, LineageContext, str]] = [] + for profile, lineage_context in zip(user_profiles, contexts, strict=True): subject_ref = self._subject_ref_for_user_id(profile.user_id) with self._lock: self._assert_subject_writable_locked(subject_ref) @@ -247,13 +264,19 @@ def add_user_profile( # out (embedding already set by precompute_profile_embeddings). if not skip_embedding: self.precompute_profile_embeddings([profile]) - embedding = profile.embedding - with self._lock: - own_txn = self._own_transaction() - try: - if own_txn: - self.conn.execute("BEGIN IMMEDIATE") + rows.append((profile, lineage_context, subject_ref)) + + with self.commit_scope(): + for profile, lineage_context, subject_ref in rows: + with self._lock: self._assert_subject_writable_locked(subject_ref) + already_exists = ( + self.conn.execute( + "SELECT 1 FROM profiles WHERE profile_id = ?", + (profile.profile_id,), + ).fetchone() + is not None + ) self.conn.execute( """INSERT OR REPLACE INTO profiles (profile_id, user_id, content, last_modified_timestamp, @@ -288,12 +311,28 @@ def add_user_profile( subject_ref, ), ) - if own_txn: - self.conn.commit() - except Exception: - if own_txn: - self.conn.rollback() - raise + if not already_exists: + _append_event_stmt( + self.conn, + org_id=self.org_id, + entity_type="profile", + entity_id=profile.profile_id, + op="create", + prov="wasGeneratedBy", + source_ids=lineage_context.source_ids, + actor=lineage_context.actor, + request_id=( + lineage_context.request_id + or profile.generated_from_request_id + or f"create_{profile.profile_id}" + ), + reason=lineage_context.reason, + model_name=lineage_context.model_name, + provider=lineage_context.provider, + ) + + for profile, _lineage_context, _subject_ref in rows: + embedding = profile.embedding fts_parts = [profile.content or ""] if profile.custom_features: fts_parts.extend(str(v) for v in profile.custom_features.values() if v) diff --git a/reflexio/server/services/storage/storage_base/__init__.py b/reflexio/server/services/storage/storage_base/__init__.py index 59d04bd51..dbe4e09a7 100644 --- a/reflexio/server/services/storage/storage_base/__init__.py +++ b/reflexio/server/services/storage/storage_base/__init__.py @@ -27,6 +27,7 @@ from ._learning_jobs import LearningJob, LearningJobStatus, LearningJobStoreABC from ._lineage import EntityType, LineageEventMixin from ._operations import OperationMixin +from ._playbook import AGGREGATE_REASON_PREFIX from ._requests import RequestMixin from ._shadow_verdicts import ShadowVerdictsMixin from ._share_links import ShareLinkMixin @@ -289,6 +290,7 @@ def learning_jobs_columns(self) -> list[str]: "PendingToolCallStatus", "PendingToolCallUpsertResult", "AgentEvaluationResultStoreMixin", + "AGGREGATE_REASON_PREFIX", "AuditEventStoreMixin", "PurgeOperationStoreMixin", "SubjectBarrierMixin", diff --git a/reflexio/server/services/storage/storage_base/playbook/_agent.py b/reflexio/server/services/storage/storage_base/playbook/_agent.py index ae74b5118..24504686c 100644 --- a/reflexio/server/services/storage/storage_base/playbook/_agent.py +++ b/reflexio/server/services/storage/storage_base/playbook/_agent.py @@ -1,6 +1,5 @@ """Abstract agent playbook CRUD + search declarations.""" -import logging from abc import abstractmethod from reflexio.models.api_schema.common import BlockingIssue @@ -9,18 +8,11 @@ PlaybookStatus, Status, ) -from reflexio.models.api_schema.domain.entities import LineageEvent +from reflexio.models.api_schema.domain.entities import LineageContext from reflexio.models.api_schema.retriever_schema import ( SearchAgentPlaybookRequest, ) from reflexio.models.config_schema import SearchOptions -from reflexio.server.tracing import capture_anomaly - -from .._playbook import AGGREGATE_REASON_PREFIX - -logger = logging.getLogger(__name__) - -_AGGREGATE_EVENT_EMIT_ATTEMPTS = 3 class AgentPlaybookStoreMixin: @@ -28,89 +20,23 @@ class AgentPlaybookStoreMixin: @abstractmethod def save_agent_playbooks( - self, agent_playbooks: list[AgentPlaybook] + self, + agent_playbooks: list[AgentPlaybook], + *, + lineage_contexts: list[LineageContext] | None = None, ) -> list[AgentPlaybook]: - """Save agent playbooks with embeddings. + """Save agent playbooks and their origin lineage events atomically. Args: agent_playbooks (list[AgentPlaybook]): List of agent playbook objects to save + lineage_contexts: Optional per-row create or aggregate attribution. + When omitted, storage emits a create event with null model/provider. Returns: list[AgentPlaybook]: Saved agent playbooks with agent_playbook_id populated from storage """ raise NotImplementedError - def save_agent_playbook_with_aggregate_event( - self, - agent_playbook: AgentPlaybook, - *, - source_ids: list[str], - request_id: str, - run_mode: str, - ) -> AgentPlaybook: - """Persist an agent playbook AND its ``op=aggregate`` lineage event. - - Backends SHOULD override this so the row insert and the event commit in ONE - transaction (the event is the sole record of the run->playbook membership for - reconstruction). This base default is a non-atomic save-then-emit fallback - with bounded retry + loud (level=error) on final failure. - - Args: - agent_playbook (AgentPlaybook): The playbook to persist. - source_ids (list[str]): IDs of the source entities that produced this playbook. - request_id (str): The aggregation run ID (used as the lineage event request_id). - run_mode (str): The aggregation run mode (e.g. ``full_archive`` or ``incremental``). - - Returns: - AgentPlaybook: The saved playbook with ``agent_playbook_id`` populated. - - Raises: - ValueError: If ``request_id`` is empty (would produce an unreconstructable event). - """ - if not request_id or not request_id.strip(): - raise ValueError( - "save_agent_playbook_with_aggregate_event requires a non-empty request_id" - ) - saved = self.save_agent_playbooks([agent_playbook])[0] - event = LineageEvent( - org_id=self.org_id, # type: ignore[attr-defined] - entity_type="agent_playbook", - entity_id=str(saved.agent_playbook_id), - op="aggregate", - prov_relation="wasDerivedFrom", - source_ids=source_ids, - actor="aggregator", - request_id=request_id, - reason=f"{AGGREGATE_REASON_PREFIX}{run_mode}", - ) - # The row is already committed; this default is non-atomic (SQLite overrides it to - # make the INSERT + event one transaction). The event is the sole reconstruction signal - # for the run, so make the emit durable: bounded retry (idempotent on retrying the - # same row's emit — entity_id is a fresh autoincrement per run, so this is NOT - # cross-run idempotency), and on final failure fail LOUD at level=error so the gap - # is paged + backfillable rather than silently lost. Never raise — the playbook - # itself is saved and must not be lost. - for attempt in range(_AGGREGATE_EVENT_EMIT_ATTEMPTS): - try: - self.append_lineage_event(event) # type: ignore[attr-defined] - return saved - except Exception: # noqa: BLE001 - logger.warning( - "aggregate lineage event append failed (attempt %d/%d) for agent_playbook %s", - attempt + 1, - _AGGREGATE_EVENT_EMIT_ATTEMPTS, - saved.agent_playbook_id, - exc_info=True, - ) - capture_anomaly( - "lineage.aggregate.append_failed", - level="error", - entity_id=str(saved.agent_playbook_id), - org_id=self.org_id, # type: ignore[attr-defined] - request_id=request_id, - ) - return saved - @abstractmethod def get_agent_playbooks( self, diff --git a/reflexio/server/services/storage/storage_base/playbook/_user.py b/reflexio/server/services/storage/storage_base/playbook/_user.py index 3c4219d50..a0899def3 100644 --- a/reflexio/server/services/storage/storage_base/playbook/_user.py +++ b/reflexio/server/services/storage/storage_base/playbook/_user.py @@ -4,6 +4,7 @@ from reflexio.models.api_schema.common import BlockingIssue from reflexio.models.api_schema.domain import Status, UserPlaybook +from reflexio.models.api_schema.domain.entities import LineageContext from reflexio.models.api_schema.retriever_schema import SearchUserPlaybookRequest from reflexio.models.config_schema import SearchOptions @@ -17,11 +18,14 @@ def save_user_playbooks( user_playbooks: list[UserPlaybook], *, skip_embedding: bool = False, + lineage_contexts: list[LineageContext] | None = None, ) -> None: - """Insert user playbooks, assigning survivor ids. + """Insert user playbooks and their create lineage events atomically. Args: user_playbooks: Playbooks to insert. + lineage_contexts: Optional per-row create attribution. When omitted, + storage derives a context from the row and leaves model/provider null. skip_embedding: When ``False`` (default — what every current caller gets), the embedding (and, when document expansion is enabled, ``expanded_terms``) is recomputed unconditionally at write time, diff --git a/reflexio/server/services/storage/storage_base/profiles/_profile_store.py b/reflexio/server/services/storage/storage_base/profiles/_profile_store.py index 80492fa33..bbaf20886 100644 --- a/reflexio/server/services/storage/storage_base/profiles/_profile_store.py +++ b/reflexio/server/services/storage/storage_base/profiles/_profile_store.py @@ -5,6 +5,7 @@ Status, UserProfile, ) +from reflexio.models.api_schema.domain.entities import LineageContext class ProfileStoreMixin: @@ -72,12 +73,15 @@ def add_user_profile( user_profiles: list[UserProfile], *, skip_embedding: bool = False, + lineage_contexts: list[LineageContext] | None = None, ) -> None: - """Add the user profile for a given user id. + """Add profiles and their create lineage events atomically. Args: user_id: The owning user id (positional, unused by some backends). user_profiles: Profiles to insert. + lineage_contexts: Optional per-row create attribution. When omitted, + storage derives a context from the row and leaves model/provider null. skip_embedding: When ``False`` (default — what every current caller gets), the embedding (and, when document expansion is enabled, ``expanded_terms``) is recomputed unconditionally at write time, diff --git a/tests/e2e_tests/test_contradiction_resolution_e2e.py b/tests/e2e_tests/test_contradiction_resolution_e2e.py index c4846c047..b467c0b05 100644 --- a/tests/e2e_tests/test_contradiction_resolution_e2e.py +++ b/tests/e2e_tests/test_contradiction_resolution_e2e.py @@ -56,6 +56,7 @@ from reflexio.models.api_schema.service_schemas import UserPlaybook from reflexio.server.api_endpoints.request_context import RequestContext +from reflexio.server.llm._litellm_types import CompletionResult, ModelProvenance from reflexio.server.llm.litellm_client import LiteLLMClient from reflexio.server.services.playbook.components.consolidator import ( DifferentiateDecision, @@ -246,8 +247,10 @@ def _drive_consolidator( tuple[list[UserPlaybook], list[int]]: ``(rows_to_save, ids_to_delete)`` as returned by ``deduplicate``. """ - consolidator.client.generate_chat_response.return_value = ( # type: ignore[attr-defined] - PlaybookConsolidationOutput(decisions=decisions) + consolidator.client.generate_chat_response_with_provenance.return_value = ( # type: ignore[attr-defined] + CompletionResult( + PlaybookConsolidationOutput(decisions=decisions), ModelProvenance() + ) ) with ( patch.object( diff --git a/tests/eval/consolidation/test_consolidation_eval.py b/tests/eval/consolidation/test_consolidation_eval.py index 42f03ade7..8b4a88d77 100644 --- a/tests/eval/consolidation/test_consolidation_eval.py +++ b/tests/eval/consolidation/test_consolidation_eval.py @@ -12,6 +12,7 @@ import pytest +from reflexio.server.llm._litellm_types import CompletionResult, ModelProvenance from reflexio.server.services.playbook.components.consolidator import ( ConsolidationDecision, DifferentiateDecision, @@ -471,8 +472,8 @@ def test_live_provider_returns_canned_decision(tmp_path): canned = UnifyDecision(new_id="NEW-0", content="x", trigger="t", rationale="r") mock = MagicMock() - mock.generate_chat_response.return_value = PlaybookConsolidationOutput( - decisions=[canned] + mock.generate_chat_response_with_provenance.return_value = CompletionResult( + PlaybookConsolidationOutput(decisions=[canned]), ModelProvenance() ) ctx = RequestContext(org_id="eval-cons-prov", storage_base_dir=str(tmp_path)) @@ -486,7 +487,7 @@ def test_live_provider_returns_canned_decision(tmp_path): assert decision == canned assert kind_for_decision(decision) == "unify" # The provider reached the LLM call (entity build + prompt render succeeded). - mock.generate_chat_response.assert_called_once() + mock.generate_chat_response_with_provenance.assert_called_once() def test_live_provider_empty_output_maps_to_independent(tmp_path): @@ -495,7 +496,9 @@ def test_live_provider_empty_output_maps_to_independent(tmp_path): from reflexio.server.api_endpoints.request_context import RequestContext mock = MagicMock() - mock.generate_chat_response.return_value = PlaybookConsolidationOutput(decisions=[]) + mock.generate_chat_response_with_provenance.return_value = CompletionResult( + PlaybookConsolidationOutput(decisions=[]), ModelProvenance() + ) ctx = RequestContext(org_id="eval-cons-noop", storage_base_dir=str(tmp_path)) provider = make_consolidation_decision_provider( diff --git a/tests/models/test_lineage_models.py b/tests/models/test_lineage_models.py index 98ff84dd0..64e251cbb 100644 --- a/tests/models/test_lineage_models.py +++ b/tests/models/test_lineage_models.py @@ -44,15 +44,44 @@ def test_lineage_event_is_content_free_and_idempotency_keyed(): actor="consolidator", request_id="req-7", reason="dup", + model_name="claude-sonnet-4-5-20250929", + provider="anthropic", ) assert e.event_id == 0 # storage assigns assert not hasattr(e, "content") + assert e.model_name == "claude-sonnet-4-5-20250929" + assert e.provider == "anthropic" + + +def test_lineage_model_provenance_defaults_to_unknown(): + event = LineageEvent( + org_id="org-42", + entity_type="profile", + entity_id="p1", + op="create", + ) + context = LineageContext(op_kind="create") + + assert event.model_name is None + assert event.provider is None + assert context.model_name is None + assert context.provider is None + assert "requested_model" not in event.model_dump() + assert "requested_model" not in context.model_dump() + assert "credential_label" not in event.model_dump() + assert "credential_label" not in context.model_dump() def test_lineage_context_and_record_ref(): ctx = LineageContext( - op_kind="merge", actor="consolidator", source_ids=["UP-1"], reason="dup" + op_kind="merge", + actor="consolidator", + source_ids=["UP-1"], + reason="dup", + model_name="claude-sonnet-4-5-20250929", + provider="anthropic", ) assert ctx.request_id is None or isinstance(ctx.request_id, str) + assert ctx.provider == "anthropic" ref = RecordRef(id="UP-2", is_purged=False) assert ref.id == "UP-2" and ref.is_purged is False diff --git a/tests/server/llm/test_claude_code_provider.py b/tests/server/llm/test_claude_code_provider.py index b1f718659..d05655a7f 100644 --- a/tests/server/llm/test_claude_code_provider.py +++ b/tests/server/llm/test_claude_code_provider.py @@ -29,6 +29,14 @@ def _stream_json(result_text: str) -> str: ) +def _stream_json_with_model(result_text: str, model: str) -> str: + return ( + json.dumps({"type": "assistant", "message": {"model": model}}) + + "\n" + + _stream_json(result_text) + ) + + @pytest.fixture(autouse=True) def _reset_module_state() -> None: """Each test starts with fresh registration and warn-once flags.""" @@ -169,12 +177,18 @@ def test_multiturn_emits_single_warning( class TestClaudeCodeLLMCompletion: def _mock_cli( - self, monkeypatch: pytest.MonkeyPatch, result_text: str = "ok" + self, + monkeypatch: pytest.MonkeyPatch, + result_text: str = "ok", + served_model: str | None = None, ) -> MagicMock: """Mock subprocess.run to return a stream-json NDJSON body with one result event.""" - mock_run = MagicMock( - return_value=_fake_completed_process(_stream_json(result_text)) + stream = ( + _stream_json_with_model(result_text, served_model) + if served_model + else _stream_json(result_text) ) + mock_run = MagicMock(return_value=_fake_completed_process(stream)) monkeypatch.setattr(ccp.subprocess, "run", mock_run) monkeypatch.setattr(ccp, "_resolve_cli_path", lambda: "/usr/local/bin/claude") return mock_run @@ -182,7 +196,11 @@ def _mock_cli( def test_basic_completion_shapes_model_response( self, monkeypatch: pytest.MonkeyPatch ) -> None: - self._mock_cli(monkeypatch, result_text="hello world") + self._mock_cli( + monkeypatch, + result_text="hello world", + served_model="claude-sonnet-5-20260701", + ) llm = ClaudeCodeLLM() response = llm.completion( @@ -192,11 +210,74 @@ def test_basic_completion_shapes_model_response( assert response.choices[0].message.content == "hello world" # type: ignore[union-attr] assert response.model == "claude-code/default" + assert ( + response._hidden_params["reflexio_served_model"] + == "claude-sonnet-5-20260701" + ) + assert response._hidden_params["reflexio_provider"] == "claude-code" + assert response._hidden_params["reflexio_cli_binary"] == "claude" # stream-json does not surface usage tokens at terminal event. assert response.usage.prompt_tokens == 0 # type: ignore[attr-defined] assert response.usage.completion_tokens == 0 # type: ignore[attr-defined] assert response.usage.total_tokens == 0 # type: ignore[attr-defined] + def test_completion_forwards_terminal_route_metadata( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + stream = ( + '{"type":"result","result":"hello","model":"MiniMax-M3",' + '"provider":"minimax"}\n' + ) + monkeypatch.setattr( + ccp.subprocess, + "run", + MagicMock(return_value=_fake_completed_process(stream)), + ) + monkeypatch.setattr(ccp, "_resolve_cli_path", lambda: "/usr/local/bin/claude") + + response = ClaudeCodeLLM().completion( + model="claude-code/default", + messages=[{"role": "user", "content": "ping"}], + ) + + assert response.model == "claude-code/default" + assert response._hidden_params["reflexio_served_model"] == "MiniMax-M3" + assert response._hidden_params["reflexio_served_provider"] == "minimax" + + def test_tool_call_response_keeps_served_model_and_binary( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + self._mock_cli( + monkeypatch, + result_text='{"tool":"finish","args":{"answer":"done"}}', + served_model="claude-sonnet-5-20260701", + ) + + response = ClaudeCodeLLM().completion( + model="claude-code/default", + messages=[{"role": "user", "content": "finish"}], + optional_params={ + "tools": [ + { + "type": "function", + "function": { + "name": "finish", + "description": "Finish", + "parameters": {"type": "object"}, + }, + } + ] + }, + ) + + assert response.model == "claude-code/default" + assert ( + response._hidden_params["reflexio_served_model"] + == "claude-sonnet-5-20260701" + ) + assert response._hidden_params["reflexio_provider"] == "claude-code" + assert response._hidden_params["reflexio_cli_binary"] == "claude" + def test_uses_stream_json_output_format( self, monkeypatch: pytest.MonkeyPatch ) -> None: @@ -559,6 +640,9 @@ def fake_run(cmd, **kwargs): assert kwargs["input"] == "Be terse.\n\n## Task\nUser: ping — now" assert kwargs["env"]["CLAUDE_SMART_HOST"] == "codex" assert response.choices[0].message.content == "codex reply" # type: ignore[union-attr] + assert response.model == "claude-code/default" + assert response._hidden_params["reflexio_provider"] == "claude-code" + assert response._hidden_params["reflexio_cli_binary"] == "codex" def test_windows_extensionless_cli_override_prefers_adjacent_cmd( self, monkeypatch: pytest.MonkeyPatch, tmp_path: Path diff --git a/tests/server/llm/test_claude_code_stream_parser.py b/tests/server/llm/test_claude_code_stream_parser.py index 90836b0d5..cd03a5dd3 100644 --- a/tests/server/llm/test_claude_code_stream_parser.py +++ b/tests/server/llm/test_claude_code_stream_parser.py @@ -22,6 +22,54 @@ def test_clean_stream_returns_success(): assert result.stall_candidate is None +def test_served_model_prefers_last_assistant_event_over_init_and_usage(): + stream = ( + '{"type":"system","subtype":"init","model":"claude-init"}\n' + '{"type":"assistant","message":{"model":"claude-first"}}\n' + '{"type":"assistant","message":{"model":"claude-served"}}\n' + '{"type":"result","result":"ok","modelUsage":{"claude-usage":{}}}\n' + ) + + result = parse_stream_json(stream, exit_code=0) + + assert result.served_model == "claude-served" + + +def test_terminal_route_metadata_is_authoritative(): + stream = ( + '{"type":"assistant","message":{"model":"bridge-model","provider":"bridge"}}\n' + '{"type":"result","result":"ok","model":"MiniMax-M3","provider":"minimax"}\n' + ) + + result = parse_stream_json(stream, exit_code=0) + + assert result.served_model == "MiniMax-M3" + assert result.served_provider == "minimax" + + +@pytest.mark.parametrize( + ("stream", "expected"), + [ + ( + '{"type":"system","subtype":"init","model":"claude-init"}\n' + '{"type":"result","result":"ok"}\n', + "claude-init", + ), + ( + '{"type":"result","result":"ok","modelUsage":{"claude-usage":{}}}\n', + "claude-usage", + ), + ( + '{"type":"result","result":"ok",' + '"modelUsage":{"claude-a":{},"claude-b":{}}}\n', + None, + ), + ], +) +def test_served_model_fallbacks_never_guess(stream, expected): + assert parse_stream_json(stream, exit_code=0).served_model == expected + + def test_billing_error_in_retry_then_stream_failure_classifies_billing(): stream = ( '{"type":"system","subtype":"api_retry","error":"billing_error","attempt":1,"max_retries":3}\n' diff --git a/tests/server/llm/test_litellm_client_surface.py b/tests/server/llm/test_litellm_client_surface.py index 0c16c4303..1cb3f7272 100644 --- a/tests/server/llm/test_litellm_client_surface.py +++ b/tests/server/llm/test_litellm_client_surface.py @@ -26,7 +26,7 @@ class must be the SAME object/class the moved code uses and tests touch — the FACADE = "reflexio.server.llm.litellm_client" -# The 5 public names (facade ``__all__``), also re-exported via server/llm/__init__. +# Public names (facade ``__all__``), also re-exported via server/llm/__init__. PUBLIC_SYMBOLS = [ "LiteLLMClient", "LiteLLMConfig", diff --git a/tests/server/llm/test_litellm_client_tool_calls.py b/tests/server/llm/test_litellm_client_tool_calls.py index cbd9f6300..854c6007a 100644 --- a/tests/server/llm/test_litellm_client_tool_calls.py +++ b/tests/server/llm/test_litellm_client_tool_calls.py @@ -6,6 +6,7 @@ import pytest +from reflexio.server.llm._litellm_types import CompletionResult from reflexio.server.llm.litellm_client import ( LiteLLMClient, LiteLLMConfig, @@ -35,6 +36,8 @@ def _mock_tool_call_response(tool_name: str, args_json: str) -> MagicMock: response = MagicMock() response.choices = [choice] response.usage = MagicMock(prompt_tokens=10, completion_tokens=5, total_tokens=15) + response.model = "gpt-4o" + response._hidden_params = {"custom_llm_provider": "openai"} return response @@ -81,7 +84,7 @@ def test_generate_chat_response_passes_tools_kwarg(self) -> None: ] with patch("litellm.completion", return_value=mock_response) as mock_completion: - result = client.generate_chat_response( + result = client.generate_chat_response_with_provenance( messages=[{"role": "user", "content": "hello"}], tools=tools, tool_choice="auto", @@ -93,11 +96,14 @@ def test_generate_chat_response_passes_tools_kwarg(self) -> None: assert call_kwargs["tool_choice"] == "auto" # The result must be a ToolCallingChatResponse - assert isinstance(result, ToolCallingChatResponse) - assert result.tool_calls is not None - assert result.tool_calls[0].function.name == "emit_profile" - assert result.finish_reason == "tool_calls" - assert result.content is None + assert isinstance(result, CompletionResult) + assert isinstance(result.value, ToolCallingChatResponse) + assert result.value.tool_calls is not None + assert result.value.tool_calls[0].function.name == "emit_profile" + assert result.value.finish_reason == "tool_calls" + assert result.value.content is None + assert result.provenance.model_name == "gpt-4o" + assert result.provenance.provider == "openai" def test_model_role_resolves_to_extraction_agent_default( self, monkeypatch: pytest.MonkeyPatch diff --git a/tests/server/llm/test_litellm_client_unit.py b/tests/server/llm/test_litellm_client_unit.py index fd3e2b308..96796e5ef 100644 --- a/tests/server/llm/test_litellm_client_unit.py +++ b/tests/server/llm/test_litellm_client_unit.py @@ -40,6 +40,8 @@ OpenAIConfig as CommonsOpenAIConfig, ) from reflexio.models.structured_output import find_schema_keyword +from reflexio.server.llm._litellm_subprocess import _snapshot_completion_response +from reflexio.server.llm._litellm_types import CompletionResult, ModelProvenance from reflexio.server.llm._provider_concurrency import ProviderCapSaturatedError from reflexio.server.llm.litellm_client import ( LiteLLMClient, @@ -461,6 +463,224 @@ def test_structured_output_pydantic(self, mock_completion): assert result.answer == "ok" assert result.score == 5 + @patch("reflexio.server.llm.litellm_client.litellm.completion") + def test_opt_in_result_carries_actual_model_and_provider(self, mock_completion): + response = _make_completion_response("hello") + response.model = "MiniMax-M3" + response._hidden_params = {"custom_llm_provider": "minimax"} + mock_completion.return_value = response + client = _build_client( + LiteLLMConfig( + model="minimax/MiniMax-M3", + api_key_config=APIKeyConfig(minimax=MiniMaxConfig(api_key="test-key")), + ) + ) + + result = client.generate_response_with_provenance("test") + + assert isinstance(result, CompletionResult) + assert result.value == "hello" + assert result.provenance == ModelProvenance( + model_name="MiniMax-M3", + provider="minimax", + ) + + @patch("reflexio.server.llm.litellm_client.litellm.completion") + def test_structured_result_provenance_does_not_serialize(self, mock_completion): + response = _make_completion_response(json.dumps({"answer": "ok", "score": 5})) + response.model = "gpt-5.4-mini" + response._hidden_params = {"custom_llm_provider": "openai"} + mock_completion.return_value = response + client = _build_client(LiteLLMConfig(model="gpt-5.4-mini")) + + result = client.generate_response_with_provenance( + "test", + response_format=SampleResponse, + ) + + assert isinstance(result, CompletionResult) + assert isinstance(result.value, SampleResponse) + assert result.value.model_dump() == {"answer": "ok", "score": 5} + assert "provenance" not in result.value.model_dump() + + @patch("reflexio.server.llm.litellm_client.litellm.completion") + def test_claude_code_does_not_launder_requested_route_as_observed( + self, mock_completion + ): + """Public ModelResponse.model is the requested route; only the stamp counts.""" + response = _make_completion_response("hello") + response.model = "claude-code/default" + response._hidden_params = { + "reflexio_provider": "claude-code", + "reflexio_cli_binary": "claude", + } + mock_completion.return_value = response + client = _build_client(LiteLLMConfig(model="claude-code/default")) + + result = client.generate_response_with_provenance("test") + + assert isinstance(result, CompletionResult) + assert result.provenance == ModelProvenance( + model_name=None, + provider=None, + ) + + @patch("reflexio.server.llm.litellm_client.litellm.completion") + def test_codex_cli_provenance_keeps_unknown_model(self, mock_completion): + response = _make_completion_response("hello") + response.model = None + response._hidden_params = { + "reflexio_provider": "claude-code", + "reflexio_cli_binary": "codex", + } + mock_completion.return_value = response + client = _build_client(LiteLLMConfig(model="claude-code/default")) + + result = client.generate_response_with_provenance("test") + + assert isinstance(result, CompletionResult) + assert result.provenance == ModelProvenance( + model_name=None, + provider=None, + ) + + @patch("reflexio.server.llm.litellm_client.litellm.completion") + def test_claude_cli_provenance_keeps_served_model(self, mock_completion): + response = _make_completion_response("hello") + response.model = "claude-sonnet-5" + response._hidden_params = { + "reflexio_provider": "claude-code", + "reflexio_cli_binary": "claude", + "reflexio_served_model": "claude-sonnet-5", + } + mock_completion.return_value = response + client = _build_client(LiteLLMConfig(model="claude-code/default")) + + result = client.generate_response_with_provenance("test") + + assert isinstance(result, CompletionResult) + assert result.provenance == ModelProvenance(model_name="claude-sonnet-5") + + @patch("reflexio.server.llm.litellm_client.litellm.completion") + def test_model_name_does_not_imply_provider(self, mock_completion): + response = _make_completion_response("hello") + response.model = "gpt-5.4-mini" + response._hidden_params = {} + mock_completion.return_value = response + client = _build_client( + LiteLLMConfig( + model="minimax/MiniMax-M3", + api_key_config=APIKeyConfig(minimax=MiniMaxConfig(api_key="test-key")), + ) + ) + + result = client.generate_response_with_provenance("test") + + assert isinstance(result, CompletionResult) + assert result.provenance == ModelProvenance(model_name="gpt-5.4-mini") + + @patch("reflexio.server.llm.litellm_client.litellm.completion") + def test_request_side_hidden_model_is_not_treated_as_served(self, mock_completion): + """LiteLLM may echo the requested model into _hidden_params['model']. + + Without a response body model or reflexio_served_model stamp, model_name + must stay unknown rather than recording the configured route as actual. + """ + response = _make_completion_response("hello") + response.model = None + response._hidden_params = { + "model": "minimax/MiniMax-M3", + "model_id": "minimax/MiniMax-M3", + "custom_llm_provider": "minimax", + } + mock_completion.return_value = response + client = _build_client( + LiteLLMConfig( + model="minimax/MiniMax-M3", + api_key_config=APIKeyConfig(minimax=MiniMaxConfig(api_key="test-key")), + ) + ) + + result = client.generate_response_with_provenance("test") + + assert isinstance(result, CompletionResult) + assert result.provenance == ModelProvenance( + model_name=None, + provider="minimax", + ) + + @patch("reflexio.server.llm.litellm_client.litellm.completion") + def test_network_fallback_uses_actual_provider(self, mock_completion): + response = _make_completion_response("hello") + response.model = "gpt-5.4-mini" + response._hidden_params = {"custom_llm_provider": "openai"} + mock_completion.return_value = response + client = _build_client( + LiteLLMConfig( + model="openai/local-model", + api_key_config=APIKeyConfig( + custom_endpoint=CustomEndpointConfig( + model="openai/local-model", + api_key="test-key", + api_base="https://example.com/v1", # type: ignore[arg-type] + ) + ), + ) + ) + + result = client.generate_response_with_provenance("test") + + assert isinstance(result, CompletionResult) + provenance = cast(CompletionResult[Any], result).provenance + assert provenance.provider == "openai" + assert provenance.model_name == "gpt-5.4-mini" + + @patch("reflexio.server.llm.litellm_client.litellm.completion") + def test_azure_provenance_uses_actual_provider(self, mock_completion): + response = _make_completion_response("hello") + response.model = "gpt-5.4-mini" + response._hidden_params = {"custom_llm_provider": "azure"} + mock_completion.return_value = response + client = _build_client( + LiteLLMConfig( + model="azure/gpt-5.4-mini", + api_key_config=APIKeyConfig( + openai=CommonsOpenAIConfig( + azure_config=AzureOpenAIConfig( + api_key="test-key", + endpoint="https://example.openai.azure.com/", # type: ignore[arg-type] + ) + ) + ), + ) + ) + + result = client.generate_response_with_provenance("test") + + assert isinstance(result, CompletionResult) + assert cast(CompletionResult[Any], result).provenance.provider == "azure" + + @patch("reflexio.server.llm.litellm_client.litellm.completion") + def test_fallback_uses_actual_provider(self, mock_completion): + response = _make_completion_response("hello") + response.model = "gpt-5.4-mini" + response._hidden_params = {"custom_llm_provider": "openai"} + mock_completion.return_value = response + client = _build_client( + LiteLLMConfig( + model="minimax/MiniMax-M3", + api_key_config=APIKeyConfig( + minimax=MiniMaxConfig(api_key="minimax-key"), + openai=CommonsOpenAIConfig(api_key="openai-key"), + ), + ) + ) + + result = client.generate_response_with_provenance("test") + + assert isinstance(result, CompletionResult) + assert cast(CompletionResult[Any], result).provenance.provider == "openai" + def test_invalid_response_format_raises(self): client = _build_client() with pytest.raises(LiteLLMClientError, match="Pydantic BaseModel class"): @@ -1722,7 +1942,11 @@ class TestStructuredOutputRepair: """Tests for opt-in corrective repair of structured output.""" def _make_mock_response( - self, content: str, *, finish_reason: str = "stop" + self, + content: str, + *, + finish_reason: str = "stop", + model: str = "served-model", ) -> MagicMock: choice = MagicMock() choice.message.content = content @@ -1730,6 +1954,7 @@ def _make_mock_response( choice.finish_reason = finish_reason resp = MagicMock() resp.choices = [choice] + resp.model = model resp.usage = MagicMock(prompt_tokens=10, completion_tokens=5, total_tokens=15) resp.usage.prompt_tokens_details = None resp.usage.cache_creation_input_tokens = None @@ -1949,6 +2174,32 @@ def fake_completion(**kwargs): assert repair_messages[1]["content"] == '{"answer": "bad", "sco' assert "truncated" in repair_messages[2]["content"] + def test_repair_error_keeps_initial_parse_failure_provenance(self): + responses = [ + ('{"answer": "bad", "sco', "served-primary"), + ('{"answer": "still bad", "score": 1}', "served-repair"), + ] + + def fake_completion(**_kwargs): + content, model = responses.pop(0) + return self._make_mock_response(content, model=model) + + client = _build_client(LiteLLMConfig(model="primary-model")) + + with ( + patch("litellm.completion", side_effect=fake_completion), + pytest.raises(StructuredOutputRepairError) as exc_info, + ): + client.generate_chat_response( + messages=[{"role": "user", "content": "test"}], + response_format=SampleResponse, + structured_output_validator=self._score_validator, + ) + + err = exc_info.value + assert err.first_parsed_provenance is not None + assert err.first_parsed_provenance.model_name == "served-repair" + def test_repair_echo_replaces_length_truncated_output(self): calls: list[dict[str, Any]] = [] responses = [ @@ -2027,6 +2278,8 @@ def fake_completion(**kwargs): ) assert not isinstance(exc_info.value, StructuredOutputRepairError) + assert exc_info.value.first_parsed_provenance is not None + assert exc_info.value.first_parsed_provenance.model_name == "served-model" assert len(calls) == 2 def test_exhaustion_keeps_latest_parsed_output_after_final_parse_failure(self): @@ -2037,9 +2290,12 @@ def test_exhaustion_keeps_latest_parsed_output_after_final_parse_failure(self): '{"answer": "first", "score": 2}', # initial: parses, semantic failure '{"answer": "esc", "sco', # repair turn: parse failure ] + served_models = ["served-primary", "served-repair"] def fake_completion(**kwargs): - return self._make_mock_response(responses.pop(0)) + return self._make_mock_response( + responses.pop(0), model=served_models.pop(0) + ) client = _build_client(LiteLLMConfig(model="primary-model")) @@ -2058,6 +2314,88 @@ def fake_completion(**kwargs): assert err.raw_content == '{"answer": "esc", "sco' assert isinstance(err.parsed_output, SampleResponse) assert err.parsed_output.score == 2 + assert err.first_parsed_provenance is not None + assert err.first_parsed_provenance.model_name == "served-primary" + + def test_ladder_preserves_first_parsed_provenance_across_rungs(self): + """Salvage attribution must match the first parse of the whole walk. + + A shared validator closure keeps the first parsed *content* across rungs. + Without ladder-wide first_parsed_provenance, the consolidator would pair + that content with the last rung's model. + """ + # Per rung: initial semantic fail + same-model repair semantic fail. + responses = [ + ('{"answer": "primary", "score": 1}', "served-primary"), + ('{"answer": "primary-repair", "score": 2}', "served-primary-repair"), + ('{"answer": "fallback", "score": 3}', "served-fallback"), + ('{"answer": "fallback-repair", "score": 4}', "served-fallback-repair"), + ] + + def fake_completion(**_kwargs): + content, model = responses.pop(0) + return self._make_mock_response(content, model=model) + + client = _build_client( + LiteLLMConfig(model="primary-model", fallback_models=["fallback-model"]) + ) + + with ( + patch("litellm.completion", side_effect=fake_completion), + pytest.raises(StructuredOutputRepairError) as exc_info, + ): + client.generate_chat_response( + messages=[{"role": "user", "content": "test"}], + response_format=SampleResponse, + structured_output_validator=self._score_validator, + ) + + err = exc_info.value + assert err.model == "fallback-model" + assert err.first_parsed_provenance is not None + assert err.first_parsed_provenance.model_name == "served-primary" + + def test_ladder_preserves_first_parsed_when_final_rung_cap_saturates(self): + """Fail-closed cap on the last rung must not drop first-parsed attribution. + + ProviderCapSaturatedError is not a LiteLLMClientError subclass. The outer + ladder must wrap it and keep ladder-wide first_parsed_provenance so + consolidator salvage pairs primary content with the primary served model. + """ + from reflexio.server.llm._provider_concurrency import ( # noqa: PLC0415 + ProviderCapSaturatedError, + ) + + call_count = 0 + + def fake_completion(**_kwargs): + nonlocal call_count + call_count += 1 + if call_count == 1: + return self._make_mock_response( + '{"answer": "primary", "score": 1}', model="served-primary" + ) + # Same-model repair + every later rung: fail-closed provider cap. + raise ProviderCapSaturatedError("provider cap saturated") + + client = _build_client( + LiteLLMConfig(model="primary-model", fallback_models=["fallback-model"]) + ) + + with ( + patch("litellm.completion", side_effect=fake_completion), + pytest.raises(LiteLLMClientError) as exc_info, + ): + client.generate_chat_response( + messages=[{"role": "user", "content": "test"}], + response_format=SampleResponse, + structured_output_validator=self._score_validator, + ) + + err = exc_info.value + assert not isinstance(err, StructuredOutputRepairError) + assert err.first_parsed_provenance is not None + assert err.first_parsed_provenance.model_name == "served-primary" # =================================================================== @@ -2950,7 +3288,9 @@ def test_per_call_max_retries_forwards_to_make_request(self, monkeypatch): monkeypatch.setattr( client, "_make_request", - lambda _messages, **kw: seen_kwargs.update(kw) or "ok", + lambda _messages, **kw: ( + seen_kwargs.update(kw) or CompletionResult("ok", ModelProvenance()) + ), ) client.generate_chat_response( [{"role": "user", "content": "hi"}], max_retries=7 @@ -2963,7 +3303,9 @@ def test_per_call_fallback_models_forwards_to_make_request(self, monkeypatch): monkeypatch.setattr( client, "_make_request", - lambda _messages, **kw: seen_kwargs.update(kw) or "ok", + lambda _messages, **kw: ( + seen_kwargs.update(kw) or CompletionResult("ok", ModelProvenance()) + ), ) client.generate_chat_response( [{"role": "user", "content": "hi"}], fallback_models=["claude-x"] @@ -2978,7 +3320,9 @@ def test_per_call_overrides_optional_default_to_config(self, monkeypatch): monkeypatch.setattr( client, "_make_request", - lambda _messages, **kw: seen_kwargs.update(kw) or "ok", + lambda _messages, **kw: ( + seen_kwargs.update(kw) or CompletionResult("ok", ModelProvenance()) + ), ) client.generate_chat_response([{"role": "user", "content": "hi"}]) assert "max_retries" not in seen_kwargs @@ -2990,6 +3334,21 @@ def test_per_call_overrides_optional_default_to_config(self, monkeypatch): # =================================================================== +def test_subprocess_snapshot_preserves_provenance_metadata(): + response = _make_completion_response("ok") + response.model = "claude-sonnet-5" + response._hidden_params = { + "reflexio_provider": "claude-code", + "reflexio_cli_binary": "claude", + "reflexio_served_model": "claude-sonnet-5", + } + + snapshot = _snapshot_completion_response(response) + + assert snapshot.model == "claude-sonnet-5" + assert snapshot._hidden_params == response._hidden_params + + class TestLitellmIntegration: """Assert _make_request hands the right knobs to litellm.completion.""" @@ -3657,6 +4016,60 @@ def _fake(**params): assert tags.get("llm.fallback_reason") == "transport_error" + def test_cli_route_resolution_is_not_reported_as_fallback(self, monkeypatch): + tags = self._install_fake_sentry(monkeypatch) + client = LiteLLMClient(LiteLLMConfig(model="claude-code/default")) + response = _make_completion_response("ok") + response.model = "claude-sonnet-5" + response._hidden_params = { + "reflexio_provider": "claude-code", + "reflexio_cli_binary": "claude", + "reflexio_served_model": "claude-sonnet-5", + } + monkeypatch.setattr("litellm.completion", lambda **_p: response) + + client.generate_chat_response([{"role": "user", "content": "hi"}]) + + assert "llm.fallback_used" not in tags + + def test_real_network_fallback_from_cli_primary_is_still_reported( + self, monkeypatch + ): + """Fallback tags fire only when the ladder advances past the primary. + + Served-model metadata on a successful CLI primary response is not a + fallback signal (see test_cli_route_resolution_is_not_reported_as_fallback). + A transport failure on the CLI primary that reaches a later rung is. + """ + tags = self._install_fake_sentry(monkeypatch) + client = LiteLLMClient( + LiteLLMConfig( + model="claude-code/default", + fallback_models=["gpt-5.4-mini"], + ) + ) + response = _make_completion_response("ok") + response.model = "gpt-5.4-mini" + response._hidden_params = {"custom_llm_provider": "openai"} + + def _fake(**params): + if params["model"] == "claude-code/default": + raise APIConnectionError( + message="cli unreachable", + llm_provider="claude-code", + model="claude-code/default", + ) + return response + + monkeypatch.setattr("litellm.completion", _fake) + + client.generate_chat_response([{"role": "user", "content": "hi"}]) + + assert tags.get("llm.fallback_used") == "true" + assert tags.get("llm.primary_model") == "claude-code/default" + assert tags.get("llm.fallback_model") == "gpt-5.4-mini" + assert tags.get("llm.fallback_reason") == "transport_error" + class TestEmbeddingRetries: """Embedding calls get num_retries parity with chat. Cross-model @@ -3762,7 +4175,8 @@ def test_generate_chat_response_does_not_mutate_caller_messages() -> None: {"role": "user", "content": "hi"}, ] - with patch.object(client, "_make_request", return_value="ok") as mock_req: + completion = CompletionResult("ok", ModelProvenance()) + with patch.object(client, "_make_request", return_value=completion) as mock_req: client.generate_chat_response(original, system_message="injected") # The caller's first dict is untouched... @@ -3773,6 +4187,6 @@ def test_generate_chat_response_does_not_mutate_caller_messages() -> None: assert sent[0]["content"] == "injected\n\norig-system" # A second call must not double-prepend onto the caller's data. - with patch.object(client, "_make_request", return_value="ok"): + with patch.object(client, "_make_request", return_value=completion): client.generate_chat_response(original, system_message="injected") assert original[0]["content"] == "orig-system" diff --git a/tests/server/llm/test_tools.py b/tests/server/llm/test_tools.py index c0114645b..9c8a5d222 100644 --- a/tests/server/llm/test_tools.py +++ b/tests/server/llm/test_tools.py @@ -6,6 +6,7 @@ import pytest from pydantic import BaseModel +from reflexio.server.llm._litellm_types import CompletionResult, ModelProvenance from reflexio.server.llm.litellm_client import ( LiteLLMClient, LiteLLMClientError, @@ -293,6 +294,44 @@ class StructuredFinish(BaseModel): assert result.structured_output.value == "ok" +def test_run_tool_loop_returns_structured_terminus_provenance(monkeypatch): + class StructuredFinish(BaseModel): + value: str + + provenance = ModelProvenance( + model_name="claude-sonnet-5", + provider="anthropic", + ) + response = CompletionResult( + ToolCallingChatResponse( + content='{"value":"ok"}', + tool_calls=None, + finish_reason="stop", + parsed_output=StructuredFinish(value="ok"), + ), + provenance, + ) + client = LiteLLMClient(LiteLLMConfig(model="claude-code/default")) + generate = MagicMock(return_value=response) + monkeypatch.setattr(client, "generate_chat_response_with_provenance", generate) + monkeypatch.setattr( + "reflexio.server.llm.tools.resolve_model_name", + lambda **_kwargs: "claude-code/default", + ) + + result = run_tool_loop( + client=client, + messages=[{"role": "user", "content": "go"}], + registry=ToolRegistry([]), + model_role=ModelRole.EXTRACTION_AGENT, + response_format=StructuredFinish, + ) + + assert result.finished_reason == "structured_output" + assert result.provenance == provenance + assert "provenance" not in result.model_dump() + + def test_run_tool_loop_records_async_accepted_and_continues( monkeypatch, tool_call_completion, @@ -382,15 +421,22 @@ def test_run_tool_loop_sends_plain_dict_tool_calls_in_followup_request(monkeypat def fake_generate_chat_response(**kwargs): calls.append(deepcopy(kwargs["messages"])) tool_calls = [emit_call] if len(calls) == 1 else [finish_call] - return ToolCallingChatResponse( - content=None, - tool_calls=tool_calls, - finish_reason="tool_calls", + return CompletionResult( + value=ToolCallingChatResponse( + content=None, + tool_calls=tool_calls, + finish_reason="tool_calls", + ), + provenance=ModelProvenance(), ) config = LiteLLMConfig(model="claude-sonnet-4-6") client = LiteLLMClient(config) - monkeypatch.setattr(client, "generate_chat_response", fake_generate_chat_response) + monkeypatch.setattr( + client, + "generate_chat_response_with_provenance", + fake_generate_chat_response, + ) ctx = LoopCtx() result = run_tool_loop( @@ -461,7 +507,11 @@ class FallbackSchema(BaseModel): emissions: list[EmitArgs] fake_parsed = FallbackSchema(emissions=[EmitArgs(value="x"), EmitArgs(value="y")]) - monkeypatch.setattr(client, "generate_chat_response", lambda **_: fake_parsed) + monkeypatch.setattr( + client, + "generate_chat_response_with_provenance", + lambda **_: CompletionResult(value=fake_parsed, provenance=ModelProvenance()), + ) ctx = LoopCtx() registry = _make_registry(ctx) @@ -503,7 +553,7 @@ def _emit_handler(args: BaseModel, c: LoopCtx) -> dict: def boom(**_kwargs): raise RuntimeError("simulated provider failure") - monkeypatch.setattr(client, "generate_chat_response", boom) + monkeypatch.setattr(client, "generate_chat_response_with_provenance", boom) result = run_tool_loop( client=client, @@ -542,7 +592,7 @@ def _emit_handler(args: BaseModel, c: LoopCtx) -> dict: def boom(**_kwargs): raise LiteLLMClientError("API call failed: hard timeout") - monkeypatch.setattr(client, "generate_chat_response", boom) + monkeypatch.setattr(client, "generate_chat_response_with_provenance", boom) with caplog.at_level(logging.WARNING, logger="reflexio.server.llm.tools"): result = run_tool_loop( @@ -694,7 +744,11 @@ def _emit(args: BaseModel, c: LoopCtx) -> dict: patch( "reflexio.server.services.service_utils.log_model_response" ) as mock_log_resp, - patch.object(client, "generate_chat_response", return_value=parsed), + patch.object( + client, + "generate_chat_response_with_provenance", + return_value=CompletionResult(value=parsed, provenance=ModelProvenance()), + ), ): run_tool_loop( client=client, @@ -759,8 +813,13 @@ def test_run_tool_loop_captures_usage_on_tool_loop_turn(monkeypatch): monkeypatch.setattr( client, - "generate_chat_response", - MagicMock(side_effect=[resp_with_usage, resp_finish]), + "generate_chat_response_with_provenance", + MagicMock( + side_effect=[ + CompletionResult(value=resp_with_usage, provenance=ModelProvenance()), + CompletionResult(value=resp_finish, provenance=ModelProvenance()), + ] + ), ) result = run_tool_loop( diff --git a/tests/server/llm/test_tools_multi_stage_integration.py b/tests/server/llm/test_tools_multi_stage_integration.py index fadeb1908..1d678220f 100644 --- a/tests/server/llm/test_tools_multi_stage_integration.py +++ b/tests/server/llm/test_tools_multi_stage_integration.py @@ -22,6 +22,7 @@ from pydantic import BaseModel, Field from reflexio.server.llm import tools as tools_mod +from reflexio.server.llm._litellm_types import CompletionResult, ModelProvenance from reflexio.server.llm.litellm_client import LiteLLMClient, LiteLLMConfig from reflexio.server.llm.model_defaults import ModelRole from reflexio.server.llm.tools import Tool, ToolRegistry, run_tool_loop @@ -114,10 +115,10 @@ def _scripted_client( client = LiteLLMClient(LiteLLMConfig(model="some-non-tool-calling-model")) iterator = iter(plans) - def fake_generate(**_kwargs: object) -> MultiStagePlan: - return next(iterator) + def fake_generate(**_kwargs: object) -> CompletionResult[MultiStagePlan]: + return CompletionResult(next(iterator), ModelProvenance()) - monkeypatch.setattr(client, "generate_chat_response", fake_generate) + monkeypatch.setattr(client, "generate_chat_response_with_provenance", fake_generate) return client diff --git a/tests/server/services/durable_learning/test_compute_persist_split.py b/tests/server/services/durable_learning/test_compute_persist_split.py index 0d21b328f..0cfc058ca 100644 --- a/tests/server/services/durable_learning/test_compute_persist_split.py +++ b/tests/server/services/durable_learning/test_compute_persist_split.py @@ -270,6 +270,7 @@ def _t() -> int: scope_enter = {"v": 0} scope_exit = {"v": 0} + scope_depth = {"v": 0} orig_scope = storage.commit_scope @@ -278,14 +279,19 @@ def wrapped_scope(): class _Tracker: def __enter__(self): - scope_enter["v"] = _t() - return cm.__enter__() + entered = cm.__enter__() + if scope_depth["v"] == 0: + scope_enter["v"] = _t() + scope_depth["v"] += 1 + return entered def __exit__(self, *exc): try: return cm.__exit__(*exc) finally: - scope_exit["v"] = _t() + scope_depth["v"] -= 1 + if scope_depth["v"] == 0: + scope_exit["v"] = _t() return _Tracker() diff --git a/tests/server/services/extraction/test_model_provenance_envelope.py b/tests/server/services/extraction/test_model_provenance_envelope.py new file mode 100644 index 000000000..2b0b0fb4d --- /dev/null +++ b/tests/server/services/extraction/test_model_provenance_envelope.py @@ -0,0 +1,90 @@ +import pytest +from pydantic import BaseModel + +from reflexio.server.llm._litellm_types import ModelProvenance +from reflexio.server.services.extraction.resumable_agent import ( + decode_committed_output, + encode_committed_output, +) + + +class _Output(BaseModel): + value: str + + +def test_committed_output_envelope_round_trips_provenance(): + provenance = ModelProvenance( + model_name="served", + provider="provider", + ) + + encoded = encode_committed_output(_Output(value="accepted"), provenance) + + output, decoded_provenance = decode_committed_output(encoded) + assert output == {"value": "accepted"} + assert decoded_provenance == provenance + + +def test_legacy_raw_output_with_colliding_keys_is_not_unwrapped(): + legacy = { + "output": {"value": "legacy"}, + "model_provenance": {"provider": "not-an-envelope"}, + } + + output, provenance = decode_committed_output(legacy) + + assert output is legacy + assert provenance is None + + +def test_v1_envelope_without_provenance_still_unwraps_output(): + output, provenance = decode_committed_output( + { + "_reflexio_envelope_version": 1, + "output": {"value": "accepted"}, + } + ) + + assert output == {"value": "accepted"} + assert provenance is None + + +@pytest.mark.parametrize( + "model_provenance", + ["not-an-object", {"provider": "provider", "unexpected": "field"}], +) +def test_v1_envelope_with_malformed_provenance_is_corrupt(model_provenance): + with pytest.raises( + ValueError, match="Corrupt v1 committed output envelope:.*model_provenance" + ): + decode_committed_output( + { + "_reflexio_envelope_version": 1, + "output": {"value": "accepted"}, + "model_provenance": model_provenance, + } + ) + + +@pytest.mark.parametrize("output", [None, "not-an-object"]) +def test_v1_envelope_with_invalid_output_is_corrupt(output): + with pytest.raises(ValueError, match="Corrupt v1 committed output envelope"): + decode_committed_output( + { + "_reflexio_envelope_version": 1, + "output": output, + } + ) + + +def test_unknown_envelope_version_is_rejected(): + future = { + "_reflexio_envelope_version": 2, + "output": {"value": "future"}, + "model_provenance": None, + } + + with pytest.raises( + ValueError, match="Unsupported committed output envelope version: 2" + ): + decode_committed_output(future) diff --git a/tests/server/services/extraction/test_resumable_agent.py b/tests/server/services/extraction/test_resumable_agent.py index e557903a9..46a2733c2 100644 --- a/tests/server/services/extraction/test_resumable_agent.py +++ b/tests/server/services/extraction/test_resumable_agent.py @@ -151,7 +151,7 @@ def test_resumable_agent_finishes_profile_output( assert stored is not None assert stored.status == AgentRunStatus.AGENT_COMPLETED assert stored.max_steps_remaining == 7 - assert stored.committed_output == { + assert stored.committed_output["output"] == { "profiles": [ { "content": "User prefers AWS ECS deployments.", @@ -162,6 +162,10 @@ def test_resumable_agent_finishes_profile_output( } ] } + assert stored.committed_output["model_provenance"] == { + "model_name": None, + "provider": None, + } def test_resumable_agent_discards_late_output_after_timeout_failure( @@ -249,7 +253,9 @@ def test_resumable_agent_finishes_playbook_output( assert stored is not None assert stored.status == AgentRunStatus.AGENT_COMPLETED assert stored.committed_output is not None - assert stored.committed_output["playbooks"][0]["trigger"] == "Deploying services" + assert stored.committed_output["output"]["playbooks"][0]["trigger"] == ( + "Deploying services" + ) def test_resumable_agent_marks_run_failed_on_loop_error(monkeypatch, storage): @@ -258,7 +264,7 @@ def test_resumable_agent_marks_run_failed_on_loop_error(monkeypatch, storage): client = LiteLLMClient(LiteLLMConfig(model="claude-sonnet-4-6")) monkeypatch.setattr( client, - "generate_chat_response", + "generate_chat_response_with_provenance", MagicMock(side_effect=RuntimeError("provider failed")), ) agent = ResumableExtractionAgent(client=client, storage=storage) @@ -296,14 +302,14 @@ def test_resumable_agent_uses_auto_tool_choice_with_extra_tools( agent = ResumableExtractionAgent(client=client, storage=storage) captured: dict[str, object] = {} - original = client.generate_chat_response + original = client.generate_chat_response_with_provenance def _spy(*args, **kwargs): captured["tool_choice"] = kwargs.get("tool_choice") captured["response_format"] = kwargs.get("response_format") return original(*args, **kwargs) - monkeypatch.setattr(client, "generate_chat_response", _spy) + monkeypatch.setattr(client, "generate_chat_response_with_provenance", _spy) extra_ctx = PendingToolCallToolContext( storage=storage, diff --git a/tests/server/services/playbook/test_aggregation_lineage_integration.py b/tests/server/services/playbook/test_aggregation_lineage_integration.py index 125c9332c..90b3cef2e 100644 --- a/tests/server/services/playbook/test_aggregation_lineage_integration.py +++ b/tests/server/services/playbook/test_aggregation_lineage_integration.py @@ -10,7 +10,7 @@ - source_ids contains str(UP-a.user_playbook_id) and str(UP-b.user_playbook_id). Also includes a regression test verifying that a failure in the atomic -``save_agent_playbook_with_aggregate_event`` ABORTS the run, restores any +``save_agent_playbooks`` ABORTS the run, restores any archived playbooks, and re-raises — all-or-nothing semantics (C1). Mirrors the real-SQLite + mocked-cluster fixture style of @@ -166,7 +166,7 @@ def test_aggregation_emits_aggregate_lineage_event( patch.object( PlaybookAggregator, "_generate_playbooks_with_source_clusters", - return_value=[(unsaved_ap, cluster_playbooks)], + return_value=[(unsaved_ap, cluster_playbooks, None)], ), ): aggregator.run(PlaybookAggregatorRequest(agent_version="v0", rerun=True)) @@ -201,7 +201,7 @@ def test_aggregate_save_failure_aborts_and_restores( aggregator: PlaybookAggregator, worker_id: str, ): - """C1: a failure in save_agent_playbook_with_aggregate_event aborts the run and restores archives. + """C1: a failure in save_agent_playbooks aborts the run and restores archives. The per-playbook save no longer silently skips on failure. Instead the exception propagates to the outer handler which: @@ -210,7 +210,7 @@ def test_aggregate_save_failure_aborts_and_restores( (c) leaves no orphan agent_playbook row (atomic rollback of the INSERT + event). Setup: seed one archived agent playbook (the old generation) + two user - playbooks. Patch save_agent_playbook_with_aggregate_event to raise. + playbooks. Patch save_agent_playbooks to raise. The archived playbook must survive (be restorable) and no new row must appear. """ org_id = request_context.org_id @@ -248,11 +248,11 @@ def test_aggregate_save_failure_aborts_and_restores( patch.object( PlaybookAggregator, "_generate_playbooks_with_source_clusters", - return_value=[(new_ap, cluster_playbooks)], + return_value=[(new_ap, cluster_playbooks, None)], ), patch.object( sqlite_storage, - "save_agent_playbook_with_aggregate_event", + "save_agent_playbooks", side_effect=RuntimeError("simulated atomic failure"), ), pytest.raises(RuntimeError, match="simulated atomic failure"), @@ -272,7 +272,7 @@ def test_aggregate_save_failure_aborts_and_restores( ap.agent_playbook_id for ap in all_aps if ap.agent_playbook_id != old_ap_id ] assert not new_ids, ( - "No new agent playbook must be saved when save_agent_playbook_with_aggregate_event fails" + "No new agent playbook must be saved when save_agent_playbooks fails" ) @@ -312,7 +312,7 @@ def test_e2e_reconstruct_added_and_run_mode( """E2E: run aggregation (full_archive) → reconstruct → assert added + run_mode. Validates the D1 rewire end-to-end: each saved playbook's aggregate event is - emitted atomically via ``save_agent_playbook_with_aggregate_event``, and + emitted atomically via ``save_agent_playbooks``, and ``reconstruct_playbook_aggregation_change_log`` can reconstruct the run with: - correct ``added_agent_playbooks`` (from aggregate events), - ``run_mode == "full_archive"`` (reason == "aggregate:full_archive"), @@ -354,7 +354,7 @@ def test_e2e_reconstruct_added_and_run_mode( patch.object( PlaybookAggregator, "_generate_playbooks_with_source_clusters", - return_value=[(new_ap, cluster_playbooks)], + return_value=[(new_ap, cluster_playbooks, None)], ), ): aggregator = PlaybookAggregator( @@ -425,7 +425,7 @@ def test_e2e_reconstruct_incremental_run_mode( patch.object( PlaybookAggregator, "_generate_playbooks_with_source_clusters", - return_value=[(ap_run1, cluster_run1)], + return_value=[(ap_run1, cluster_run1, None)], ), ): agg1 = PlaybookAggregator( @@ -455,7 +455,7 @@ def test_e2e_reconstruct_incremental_run_mode( patch.object( PlaybookAggregator, "_generate_playbooks_with_source_clusters", - return_value=[(ap_run2, cluster_run2)], + return_value=[(ap_run2, cluster_run2, None)], ), ): agg2 = PlaybookAggregator( diff --git a/tests/server/services/playbook/test_aggregation_soft_delete_integration.py b/tests/server/services/playbook/test_aggregation_soft_delete_integration.py index 37a491767..73fd515b6 100644 --- a/tests/server/services/playbook/test_aggregation_soft_delete_integration.py +++ b/tests/server/services/playbook/test_aggregation_soft_delete_integration.py @@ -366,7 +366,7 @@ def _run_aggregator_with_supersede( patch.object( PlaybookAggregator, "_generate_playbooks_with_source_clusters", - return_value=[(new_ap, cluster_playbooks)], + return_value=[(new_ap, cluster_playbooks, None)], ), ): aggregator = PlaybookAggregator( @@ -601,7 +601,7 @@ def __str__(self) -> str: patch.object( PlaybookAggregator, "_generate_playbooks_with_source_clusters", - return_value=[(new_ap, cluster_playbooks)], + return_value=[(new_ap, cluster_playbooks, None)], ), patch(uuid_path, return_value=_EmptyStrUUID()), ): @@ -611,7 +611,7 @@ def __str__(self) -> str: agent_version="v0", ) # The storage guard raises on empty request_id; outer handler restores + re-raises. - with pytest.raises(StorageError, match="non-empty request_id"): + with pytest.raises(StorageError, match="requires request_id"): aggregator.run( PlaybookAggregatorRequest(agent_version="v0", rerun=True) ) diff --git a/tests/server/services/playbook/test_apply_playbook_edit_integration.py b/tests/server/services/playbook/test_apply_playbook_edit_integration.py index f38bda1e2..d98025f5f 100644 --- a/tests/server/services/playbook/test_apply_playbook_edit_integration.py +++ b/tests/server/services/playbook/test_apply_playbook_edit_integration.py @@ -53,7 +53,7 @@ def test_apply_no_orphan_when_incumbent_already_gone(tmp_path): assert all(p.content != "v2" for p in currents) -def test_apply_lost_cas_deletes_inserted_successor_and_leaves_no_orphan(tmp_path): +def test_apply_lost_cas_rolls_back_successor_and_leaves_no_orphan(tmp_path): s = SQLiteStorage(org_id="test_org", db_path=str(tmp_path / "t.db")) s.migrate() inc = UserPlaybook(user_id="u", agent_version="v", request_id="r", content="v1") @@ -90,4 +90,4 @@ def test_apply_lost_cas_deletes_inserted_successor_and_leaves_no_orphan(tmp_path assert len(currents) == 1 assert currents[0].content == "winner" events = s.get_lineage_events(entity_type="user_playbook") - assert [e.op for e in events] == ["revise"] + assert [e.op for e in events] == ["create", "create", "revise"] diff --git a/tests/server/services/playbook/test_cluster_change_detection.py b/tests/server/services/playbook/test_cluster_change_detection.py index 50f1a0d30..1e1765c1e 100644 --- a/tests/server/services/playbook/test_cluster_change_detection.py +++ b/tests/server/services/playbook/test_cluster_change_detection.py @@ -24,6 +24,7 @@ def disable_mock_llm_response(monkeypatch): UserPlaybook, ) from reflexio.models.config_schema import PlaybookAggregatorConfig +from reflexio.server.llm._litellm_types import CompletionResult, ModelProvenance from reflexio.server.services.playbook.components.aggregator import ( PlaybookAggregator, ) @@ -418,7 +419,9 @@ def _setup_aggregator_for_run( mock_storage.get_user_playbooks.return_value = user_playbooks mock_storage.get_agent_playbooks.return_value = existing_playbooks mock_storage.count_user_playbooks.return_value = len(user_playbooks) - mock_storage.save_agent_playbooks.return_value = [] + mock_storage.save_agent_playbooks.side_effect = lambda playbooks, **_kwargs: ( + playbooks + ) # Setup operation state (for fingerprints and bookmarks) # Storage returns {"operation_state": {...}} wrapping @@ -438,7 +441,9 @@ def get_operation_state_side_effect(key): trigger="When something happens", ) mock_response = PlaybookAggregationOutput(playbook=structured) - mock_llm_client.generate_chat_response.return_value = mock_response + mock_llm_client.generate_chat_response_with_provenance.return_value = ( + CompletionResult(mock_response, ModelProvenance()) + ) mock_llm_client.config = MagicMock() mock_llm_client.config.model = "test-model" @@ -461,17 +466,16 @@ def test_first_run_calls_llm_for_all_clusters(self): operation_state=None, ) - # Make save_agent_playbook_with_aggregate_event return playbooks with IDs + # Make save_agent_playbooks return playbooks with IDs _id_counter = [0] - def save_with_event_side_effect(playbook, *, source_ids, request_id, run_mode): # noqa: ANN001, ARG001 + def save_with_event_side_effect(playbooks, **_kwargs): # noqa: ANN001 _id_counter[0] += 1 + playbook = playbooks[0] playbook.agent_playbook_id = _id_counter[0] - return playbook + return [playbook] - mock_storage.save_agent_playbook_with_aggregate_event.side_effect = ( - save_with_event_side_effect - ) + mock_storage.save_agent_playbooks.side_effect = save_with_event_side_effect request = PlaybookAggregatorRequest( agent_version="1.0", @@ -480,9 +484,9 @@ def save_with_event_side_effect(playbook, *, source_ids, request_id, run_mode): aggregator.run(request) # LLM should be called for each cluster (at least 1, up to 2) - assert mock_llm_client.generate_chat_response.call_count >= 1 - # save_agent_playbook_with_aggregate_event should be called - mock_storage.save_agent_playbook_with_aggregate_event.assert_called() + assert mock_llm_client.generate_chat_response_with_provenance.call_count >= 1 + # save_agent_playbooks should be called + mock_storage.save_agent_playbooks.assert_called() # Fingerprints should be stored mock_storage.upsert_operation_state.assert_called() @@ -584,14 +588,13 @@ def test_second_run_with_new_playbooks_calls_llm_selectively(self): _id_counter2 = [200] - def save_with_event_side_effect2(playbook, *, source_ids, request_id, run_mode): # noqa: ANN001, ARG001 + def save_with_event_side_effect2(playbooks, **_kwargs): # noqa: ANN001 _id_counter2[0] += 1 + playbook = playbooks[0] playbook.agent_playbook_id = _id_counter2[0] - return playbook + return [playbook] - mock_storage.save_agent_playbook_with_aggregate_event.side_effect = ( - save_with_event_side_effect2 - ) + mock_storage.save_agent_playbooks.side_effect = save_with_event_side_effect2 request = PlaybookAggregatorRequest( agent_version="1.0", @@ -600,10 +603,12 @@ def save_with_event_side_effect2(playbook, *, source_ids, request_id, run_mode): aggregator.run(request) # LLM should be called fewer times than total clusters - total_llm_calls = mock_llm_client.generate_chat_response.call_count + total_llm_calls = ( + mock_llm_client.generate_chat_response_with_provenance.call_count + ) assert total_llm_calls >= 1 - # save_agent_playbook_with_aggregate_event should be called - mock_storage.save_agent_playbook_with_aggregate_event.assert_called() + # save_agent_playbooks should be called + mock_storage.save_agent_playbooks.assert_called() def test_rerun_bypasses_change_detection(self): """rerun=True should call LLM for ALL clusters regardless of fingerprints.""" @@ -640,14 +645,13 @@ def test_rerun_bypasses_change_detection(self): _id_counter3 = [0] - def save_with_event_side_effect3(playbook, *, source_ids, request_id, run_mode): # noqa: ANN001, ARG001 + def save_with_event_side_effect3(playbooks, **_kwargs): # noqa: ANN001 _id_counter3[0] += 1 + playbook = playbooks[0] playbook.agent_playbook_id = _id_counter3[0] - return playbook + return [playbook] - mock_storage.save_agent_playbook_with_aggregate_event.side_effect = ( - save_with_event_side_effect3 - ) + mock_storage.save_agent_playbooks.side_effect = save_with_event_side_effect3 request = PlaybookAggregatorRequest( agent_version="1.0", @@ -657,7 +661,9 @@ def save_with_event_side_effect3(playbook, *, source_ids, request_id, run_mode): aggregator.run(request) # LLM should be called for ALL clusters - assert mock_llm_client.generate_chat_response.call_count == len(clusters) + assert mock_llm_client.generate_chat_response_with_provenance.call_count == len( + clusters + ) # archive_agent_playbooks_by_playbook_name should be called for each # full-archive playbook name (one call per name) mock_storage.archive_agent_playbooks_by_playbook_name.assert_called() @@ -742,14 +748,13 @@ def test_first_run_supersedes_archived_on_success(self): _id_counter4 = [0] - def save_with_event_side_effect4(playbook, *, source_ids, request_id, run_mode): # noqa: ANN001, ARG001 + def save_with_event_side_effect4(playbooks, **_kwargs): # noqa: ANN001 _id_counter4[0] += 1 + playbook = playbooks[0] playbook.agent_playbook_id = _id_counter4[0] - return playbook + return [playbook] - mock_storage.save_agent_playbook_with_aggregate_event.side_effect = ( - save_with_event_side_effect4 - ) + mock_storage.save_agent_playbooks.side_effect = save_with_event_side_effect4 request = PlaybookAggregatorRequest( agent_version="1.0", @@ -805,7 +810,9 @@ def test_raw_string_response_returns_none(self): mock_request_context.configurator = MagicMock() # LLM returns a raw string instead of PlaybookAggregationOutput - mock_llm_client.generate_chat_response.return_value = "unparsed text" + mock_llm_client.generate_chat_response_with_provenance.return_value = ( + CompletionResult("unparsed text", ModelProvenance()) + ) mock_llm_client.config = MagicMock() mock_llm_client.config.model = "test-model" @@ -841,8 +848,10 @@ def test_valid_aggregation_output_is_processed(self): content="Be concise when answering questions", trigger="When answering questions", ) - mock_llm_client.generate_chat_response.return_value = PlaybookAggregationOutput( - playbook=structured + mock_llm_client.generate_chat_response_with_provenance.return_value = ( + CompletionResult( + PlaybookAggregationOutput(playbook=structured), ModelProvenance() + ) ) mock_llm_client.config = MagicMock() mock_llm_client.config.model = "test-model" @@ -867,9 +876,11 @@ def test_valid_aggregation_output_is_processed(self): result = aggregator._generate_playbook_from_cluster(cluster_playbooks, "None") assert result is not None - assert result.content == "Be concise when answering questions" - assert result.trigger == "When answering questions" - assert result.playbook_status == PlaybookStatus.PENDING + playbook, provenance = result + assert playbook.content == "Be concise when answering questions" + assert playbook.trigger == "When answering questions" + assert playbook.playbook_status == PlaybookStatus.PENDING + assert provenance == ModelProvenance() class TestClusteringStability: diff --git a/tests/server/services/playbook/test_consolidation_lineage_integration.py b/tests/server/services/playbook/test_consolidation_lineage_integration.py index 7f480de06..742c72ec2 100644 --- a/tests/server/services/playbook/test_consolidation_lineage_integration.py +++ b/tests/server/services/playbook/test_consolidation_lineage_integration.py @@ -25,10 +25,12 @@ from reflexio.models.api_schema.domain.enums import Status from reflexio.models.api_schema.service_schemas import UserPlaybook from reflexio.server.api_endpoints.request_context import RequestContext +from reflexio.server.llm._litellm_types import CompletionResult, ModelProvenance from reflexio.server.services.lineage.resolve import resolve_current from reflexio.server.services.playbook.components.consolidator import ( DifferentiateDecision, PlaybookConsolidationOutput, + PlaybookConsolidator, UnifyDecision, ) from reflexio.server.services.playbook.service import ( @@ -173,6 +175,19 @@ def test_consolidation_merge_routes_through_merge_records( ] ) + observed_provenance = ModelProvenance( + model_name="MiniMax-M3", + provider="minimax", + ) + real_deduplicate = PlaybookConsolidator.deduplicate + + def _deduplicate_with_observed_provenance(self, *args, **kwargs): + result = real_deduplicate(self, *args, **kwargs) + # Stamp observed consolidator attribution so the merge event path is + # exercised without a live LLM completion. + self.model_provenance = observed_provenance + return result + with ( patch.object( PlaybookGenerationService, @@ -187,6 +202,11 @@ def test_consolidation_merge_routes_through_merge_records( "reflexio.server.services.playbook.components.consolidator.PlaybookConsolidator._consolidation_decisions", return_value=decision_output, ), + patch.object( + PlaybookConsolidator, + "deduplicate", + _deduplicate_with_observed_provenance, + ), patch.dict("os.environ", {"MOCK_LLM_RESPONSE": "false"}), ): generation_service._finalize_extracted_items([_candidate()]) @@ -205,13 +225,15 @@ def test_consolidation_merge_routes_through_merge_records( assert tombstone.status == Status.MERGED assert tombstone.merged_into == survivor.user_playbook_id - # A merge lineage event keyed on the survivor exists. + # A merge lineage event keyed on the survivor exists, with consolidator model. events = sqlite_storage.get_lineage_events( entity_type="user_playbook", entity_id=str(survivor.user_playbook_id) ) merge_events = [e for e in events if e.op == "merge"] assert len(merge_events) == 1, events assert str(old_id) in merge_events[0].source_ids + assert merge_events[0].model_name == observed_provenance.model_name + assert merge_events[0].provider == observed_provenance.provider # resolve_current follows merged_into to the live survivor. ref = resolve_current(sqlite_storage, "user_playbook", old_id) @@ -263,9 +285,9 @@ def test_consolidation_repair_persists_only_repaired_multi_new_unify( def repaired_consolidation(*, structured_output_validator, **_kwargs): assert structured_output_validator(initial_output) assert structured_output_validator(repaired_output) == [] - return repaired_output + return CompletionResult(repaired_output, ModelProvenance()) - generation_service.client.generate_chat_response.side_effect = ( + generation_service.client.generate_chat_response_with_provenance.side_effect = ( repaired_consolidation ) @@ -288,7 +310,7 @@ def repaired_consolidation(*, structured_output_validator, **_kwargs): survivor = current[0] assert survivor.content == "Always update target groups and security groups." assert survivor.source_interaction_ids == [15, 16, 19, 20] - generation_service.client.generate_chat_response.assert_called_once() + generation_service.client.generate_chat_response_with_provenance.assert_called_once() def test_consolidation_differentiate_tombstones_split_source( diff --git a/tests/server/services/playbook/test_extractor_polarity_integration.py b/tests/server/services/playbook/test_extractor_polarity_integration.py index ffe7770b5..c70a15fcd 100644 --- a/tests/server/services/playbook/test_extractor_polarity_integration.py +++ b/tests/server/services/playbook/test_extractor_polarity_integration.py @@ -27,6 +27,7 @@ ) from reflexio.models.config_schema import PlaybookConfig from reflexio.server.api_endpoints.request_context import RequestContext +from reflexio.server.llm._litellm_types import CompletionResult, ModelProvenance from reflexio.server.llm.litellm_client import LiteLLMClient from reflexio.server.services.playbook.components.extractor import PlaybookExtractor from reflexio.server.services.playbook.playbook_service_utils import ( @@ -86,6 +87,15 @@ def mock_llm_client(): client = MagicMock(spec=LiteLLMClient) client.config = LiteLLMConfig(model="claude-sonnet-4-6") + + def _generate_with_provenance(*args, **kwargs): + return CompletionResult( + client.generate_chat_response(*args, **kwargs), ModelProvenance() + ) + + client.generate_chat_response_with_provenance.side_effect = ( + _generate_with_provenance + ) return client diff --git a/tests/server/services/playbook/test_playbook_aggregator.py b/tests/server/services/playbook/test_playbook_aggregator.py index 17b857967..5af287b0c 100644 --- a/tests/server/services/playbook/test_playbook_aggregator.py +++ b/tests/server/services/playbook/test_playbook_aggregator.py @@ -30,6 +30,7 @@ PlaybookConfig, PlaybookOptimizerConfig, ) +from reflexio.server.llm._litellm_types import CompletionResult, ModelProvenance from reflexio.server.services.playbook.aggregation_prompt_processing import ( AggregationPromptProcessingContext, PromptPostprocessResult, @@ -54,6 +55,14 @@ def _make_aggregator( ) -> Any: """Build an aggregator with fully mocked dependencies.""" llm = MagicMock() + + def _generate_with_provenance(*args: Any, **kwargs: Any) -> CompletionResult[Any]: + value = llm.generate_chat_response(*args, **kwargs) + if isinstance(value, CompletionResult): + return value + return CompletionResult(value=value, provenance=ModelProvenance()) + + llm.generate_chat_response_with_provenance.side_effect = _generate_with_provenance ctx = MagicMock() ctx.storage = storage or MagicMock() ctx.configurator = configurator or MagicMock() @@ -378,7 +387,7 @@ def test_mock_llm_response_postprocesses_artifacts_before_storage(): result = agg._generate_playbooks_with_source_clusters(clusters, []) assert len(result) == 1 - playbook, _sources = result[0] + playbook, _sources, _provenance = result[0] assert "< the do-rule and the avoid-rule were not collapsed. bullet_lines = [ - line for line in result.content.splitlines() if line.strip().startswith("-") + line for line in playbook.content.splitlines() if line.strip().startswith("-") ] assert len(bullet_lines) == 2 @@ -2138,11 +2196,12 @@ def _make_ordered_aggregator(call_log: list[tuple[str, Any]]) -> Any: agg.storage.get_agent_playbooks.return_value = [] agg.storage.get_user_playbooks.return_value = [_raw(rid=1), _raw(rid=2)] - def _save(playbook: AgentPlaybook, **_kwargs: Any) -> AgentPlaybook: + def _save(playbooks: list[AgentPlaybook], **_kwargs: Any) -> list[AgentPlaybook]: + playbook = playbooks[0] call_log.append(("save", playbook.agent_playbook_id)) - return playbook + return [playbook] - agg.storage.save_agent_playbook_with_aggregate_event.side_effect = _save + agg.storage.save_agent_playbooks.side_effect = _save def _set_source_windows(agent_playbook_id: int, _windows: Any) -> None: call_log.append(("set_source_windows", agent_playbook_id)) @@ -2180,7 +2239,9 @@ def _restore_by_name(name: str, **_kwargs: Any) -> None: def _instrument_run( call_log: list[tuple[str, Any]], clusters: dict[int, list[UserPlaybook]], - generated_pairs: list[tuple[AgentPlaybook, list[UserPlaybook]]], + generated_pairs: list[ + tuple[AgentPlaybook, list[UserPlaybook], ModelProvenance | None] + ], *, uuid_side_effect: Any | None = None, ): @@ -2254,7 +2315,7 @@ def _two_pairs( pb_a.agent_playbook_id = 100 pb_b = _agent_playbook(fid=200) pb_b.agent_playbook_id = 200 - generated_pairs = [(pb_a, cluster_a), (pb_b, cluster_b)] + generated_pairs = [(pb_a, cluster_a, None), (pb_b, cluster_b, None)] return clusters, generated_pairs def test_archive_between_generate_and_first_save(self): diff --git a/tests/server/services/playbook/test_playbook_consolidator.py b/tests/server/services/playbook/test_playbook_consolidator.py index 070f68acd..c06826711 100644 --- a/tests/server/services/playbook/test_playbook_consolidator.py +++ b/tests/server/services/playbook/test_playbook_consolidator.py @@ -11,7 +11,11 @@ import pytest from reflexio.models.api_schema.service_schemas import UserPlaybook -from reflexio.server.llm.litellm_client import StructuredOutputRepairError +from reflexio.server.llm._litellm_types import CompletionResult, ModelProvenance +from reflexio.server.llm.litellm_client import ( + LiteLLMClientError, + StructuredOutputRepairError, +) from reflexio.server.services.playbook.components.consolidator import ( DifferentiateDecision, IndependentDecision, @@ -58,6 +62,16 @@ def mock_consolidator(): mock_llm_client = MagicMock() + def _generate_with_provenance(*args, **kwargs): + value = mock_llm_client.generate_chat_response(*args, **kwargs) + if isinstance(value, CompletionResult): + return value + return CompletionResult(value=value, provenance=ModelProvenance()) + + mock_llm_client.generate_chat_response_with_provenance.side_effect = ( + _generate_with_provenance + ) + with patch( "reflexio.server.services.deduplication_utils.SiteVarManager" ) as mock_svm: @@ -95,6 +109,11 @@ def _unify( def _shared_repair_side_effect(*outputs: PlaybookConsolidationOutput): """Simulate the shared client validator/repair contract for service tests.""" + provenance = ModelProvenance( + model_name="served-model", + provider="provider", + ) + def _side_effect(*, response_format, structured_output_validator, model, **_kwargs): assert response_format is PlaybookConsolidationOutput first_output = outputs[0] @@ -112,6 +131,7 @@ def _side_effect(*, response_format, structured_output_validator, model, **_kwar model=model, parsed_output=repaired_output, validation_errors=tuple(repaired_errors), + first_parsed_provenance=provenance, ) raise StructuredOutputRepairError( "repair exhausted", @@ -119,6 +139,7 @@ def _side_effect(*, response_format, structured_output_validator, model, **_kwar model=model, parsed_output=first_output, validation_errors=tuple(first_errors), + first_parsed_provenance=provenance, ) return _side_effect @@ -1180,6 +1201,78 @@ def test_unify_against_existing_archives_the_existing(self, mock_consolidator): class TestConsolidationRepair: """Tests for pre-apply validation and the single repair pass.""" + def test_repair_fallback_uses_first_parsed_output_provenance( + self, mock_consolidator + ): + first_output = PlaybookConsolidationOutput( + decisions=[IndependentDecision(new_id="NEW-0")] + ) + first_parsed_provenance = ModelProvenance(model_name="served-first-parsed") + + def repair_exhausted(*, structured_output_validator, **_kwargs): + structured_output_validator(first_output) + raise StructuredOutputRepairError( + "repair exhausted", + failure_kind="semantic", + model="configured-model", + first_parsed_provenance=first_parsed_provenance, + ) + + mock_consolidator.client.generate_chat_response.side_effect = repair_exhausted + + result = mock_consolidator._consolidation_decisions( + [_make_user_playbook(0), _make_user_playbook(1)], [] + ) + + assert result is first_output + assert mock_consolidator.model_provenance == first_parsed_provenance + + def test_repair_fallback_returns_first_parsed_output_without_provenance( + self, mock_consolidator + ): + first_output = PlaybookConsolidationOutput( + decisions=[IndependentDecision(new_id="NEW-0")] + ) + + def repair_exhausted(*, structured_output_validator, **_kwargs): + structured_output_validator(first_output) + raise StructuredOutputRepairError( + "repair exhausted", + failure_kind="semantic", + model="configured-model", + ) + + mock_consolidator.client.generate_chat_response.side_effect = repair_exhausted + + result = mock_consolidator._consolidation_decisions( + [_make_user_playbook(0), _make_user_playbook(1)], [] + ) + + assert result is first_output + assert mock_consolidator.model_provenance is None + + def test_transport_failure_preserves_first_parsed_output(self, mock_consolidator): + first_output = PlaybookConsolidationOutput( + decisions=[IndependentDecision(new_id="NEW-0")] + ) + first_parsed_provenance = ModelProvenance(model_name="served-first-parsed") + + def transport_failure(*, structured_output_validator, **_kwargs): + structured_output_validator(first_output) + raise LiteLLMClientError( + "repair transport failed", + first_parsed_provenance=first_parsed_provenance, + ) + + mock_consolidator.client.generate_chat_response.side_effect = transport_failure + + result = mock_consolidator._consolidation_decisions( + [_make_user_playbook(0), _make_user_playbook(1)], [] + ) + + assert result is first_output + assert mock_consolidator.model_provenance == first_parsed_provenance + def test_under_consumed_output_repairs_to_multi_new_unify(self, mock_consolidator): new_0 = _make_user_playbook( 0, content="alpha beta", source_interaction_ids=[10] @@ -1278,6 +1371,8 @@ def test_repair_failure_falls_back_to_original_output( ] assert delete_ids == [] assert mock_consolidator.client.generate_chat_response.call_count == 1 + assert mock_consolidator.model_provenance is not None + assert mock_consolidator.model_provenance.model_name == "served-model" def test_suspicious_same_source_split_triggers_repair(self, mock_consolidator): new_0 = _make_user_playbook( diff --git a/tests/server/services/playbook/test_playbook_consolidator_integration.py b/tests/server/services/playbook/test_playbook_consolidator_integration.py index f0a41b8ba..d39341663 100644 --- a/tests/server/services/playbook/test_playbook_consolidator_integration.py +++ b/tests/server/services/playbook/test_playbook_consolidator_integration.py @@ -30,6 +30,7 @@ from reflexio.models.api_schema.service_schemas import UserPlaybook from reflexio.server.api_endpoints.request_context import RequestContext +from reflexio.server.llm._litellm_types import CompletionResult, ModelProvenance from reflexio.server.llm.litellm_client import LiteLLMClient, LiteLLMConfig from reflexio.server.services.playbook.components.consolidator import ( DifferentiateDecision, @@ -81,7 +82,18 @@ def request_context(sqlite_storage, temp_storage_dir, worker_id): @pytest.fixture def mock_llm_client(): """Mock LiteLLM client. ``generate_chat_response`` is set per-test.""" - return MagicMock(spec=LiteLLMClient) + client = MagicMock(spec=LiteLLMClient) + + def _generate_with_provenance(*args, **kwargs): + value = client.generate_chat_response(*args, **kwargs) + if isinstance(value, CompletionResult): + return value + return CompletionResult(value, ModelProvenance()) + + client.generate_chat_response_with_provenance.side_effect = ( + _generate_with_provenance + ) + return client @pytest.fixture @@ -972,9 +984,7 @@ class TestConsolidatorNativeFallbackEndToEnd: every rung and NO ``fallbacks`` kwarg ever handed to litellm. """ - def test_consolidator_advances_to_fallback_rung( - self, request_context, monkeypatch - ): + def test_consolidator_advances_to_fallback_rung(self, request_context, monkeypatch): """Production-style: ``REFLEXIO_LLM_FALLBACK_MODELS`` set globally. When the primary fails, the owned walk advances to the configured diff --git a/tests/server/services/playbook/test_playbook_edit_apply.py b/tests/server/services/playbook/test_playbook_edit_apply.py index 2872cc8f6..cee4198e8 100644 --- a/tests/server/services/playbook/test_playbook_edit_apply.py +++ b/tests/server/services/playbook/test_playbook_edit_apply.py @@ -120,8 +120,8 @@ def test_apply_expect_current_false_archives(): def test_apply_expect_current_false_returns_minus1_and_no_orphan(): """When incumbent is already archived, supersede_record returns False. - The new code deletes the just-inserted successor so no orphan CURRENT row - remains — the -1 return value indicates the lost race, not an orphan. + The transaction rolls back the provisional successor and its create event, + so the -1 return value indicates the lost race, not an orphan. """ from reflexio.server.services.playbook.playbook_edit_apply import ( apply_playbook_edit, @@ -137,6 +137,9 @@ def test_apply_expect_current_false_returns_minus1_and_no_orphan(): # Archive first so supersede_record will return False s.archive_user_playbook_by_id(user_id="u1", user_playbook_id=old_id) + event_ids_before = { + event.event_id for event in s.get_lineage_events(org_id="org_apply_1") + } new = _playbook(content="new") new_id = apply_playbook_edit( @@ -146,13 +149,16 @@ def test_apply_expect_current_false_returns_minus1_and_no_orphan(): source="offline_optimizer", request_id="run-abc", ) - # supersede_record returned False → -1, successor cleaned up (no orphan) + # supersede_record returned False → -1, transaction rolled back. assert new_id == -1 # No orphan: the inserted successor was deleted all_pbs = s.get_user_playbooks(user_id="u1") current_ids = {p.user_playbook_id for p in all_pbs if p.status is None} assert len(current_ids) == 0 + assert { + event.event_id for event in s.get_lineage_events(org_id="org_apply_1") + } == event_ids_before def test_apply_raises_on_empty_request_id_before_write(): @@ -249,9 +255,9 @@ def test_apply_lineage_event_carries_operation_run_id(): events = s.get_lineage_events( entity_type="user_playbook", entity_id=str(new_id) ) - assert len(events) == 1 - assert events[0].op == "revise" - assert events[0].request_id == operation_run_id, ( + assert [event.op for event in events] == ["create", "revise"] + revise_event = events[1] + assert revise_event.request_id == operation_run_id, ( f"lineage event must carry the operation run id {operation_run_id!r}, " f"not the incumbent's birth request_id {old.request_id!r}" ) diff --git a/tests/server/services/playbook/test_playbook_generation_service.py b/tests/server/services/playbook/test_playbook_generation_service.py index bbba1c71f..905edb83f 100644 --- a/tests/server/services/playbook/test_playbook_generation_service.py +++ b/tests/server/services/playbook/test_playbook_generation_service.py @@ -23,6 +23,7 @@ ) from reflexio.server.api_endpoints.request_context import RequestContext from reflexio.server.extensions import register_service +from reflexio.server.llm._litellm_types import ModelProvenance from reflexio.server.llm.litellm_client import LiteLLMClient, LiteLLMConfig from reflexio.server.services.playbook.aggregation_prompt_processing import ( AGGREGATION_PROMPT_PROCESSOR, @@ -628,7 +629,7 @@ def test_error_handling(mock_chat_completion): auto_run=False, ) - # Mock storage.save_user_playbooks to raise an exception + # Mock the lineage-aware save to raise an exception with patch.object( _storage(playbook_generation_service), "save_user_playbooks", @@ -684,7 +685,8 @@ def test_finalize_drops_empty_and_same_batch_duplicates_with_dedup_flag_off(): "reflexio.server.services.playbook.components.consolidator.PlaybookConsolidator", ) as mock_dedup_cls, patch.object( - _storage(playbook_generation_service), "save_user_playbooks" + _storage(playbook_generation_service), + "save_user_playbooks", ) as save_user_playbooks, patch.object( playbook_generation_service, "_enqueue_user_playbook_optimization" @@ -698,7 +700,11 @@ def test_finalize_drops_empty_and_same_batch_duplicates_with_dedup_flag_off(): ) ) playbook_generation_service._finalize_extracted_items( - [first, duplicate, blank] + [first, duplicate, blank], + model_provenance=ModelProvenance( + model_name="served-model", + provider="provider", + ), ) save_user_playbooks.assert_called_once() @@ -706,6 +712,58 @@ def test_finalize_drops_empty_and_same_batch_duplicates_with_dedup_flag_off(): assert saved_playbooks == [first] assert first.status is None assert first.source == "test_source" + context = save_user_playbooks.call_args.kwargs["lineage_contexts"][0] + assert context.op_kind == "create" + assert context.model_name == "served-model" + + +def test_finalize_without_provenance_emits_create_with_null_model_fields(): + """Opaque routes still write create lineage; model fields stay null.""" + from reflexio.models.api_schema.domain.entities import LineageContext + + with tempfile.TemporaryDirectory() as temp_dir: + service = PlaybookGenerationService( + llm_client=LiteLLMClient(LiteLLMConfig(model="gpt-4o-mini")), + request_context=RequestContext(org_id="0", storage_base_dir=temp_dir), + ) + service.service_config = PlaybookGenerationServiceConfig( + request_id="legacy-request", + agent_version="1.0", + user_id="test-user", + source="test", + ) + playbook = UserPlaybook( + agent_version="1.0", + request_id="legacy-request", + content="Preserve the old output.", + trigger="When resuming a legacy run", + ) + + with ( + patch( + "reflexio.server.services.playbook.components.consolidator.PlaybookConsolidator", + ) as mock_dedup_cls, + patch.object(_storage(service), "save_user_playbooks") as save, + patch.object(service, "_enqueue_user_playbook_optimization"), + ): + mock_dedup_cls.return_value.deduplicate.return_value = ( + [playbook], + [], + [], + ) + mock_dedup_cls.return_value.model_provenance = None + mock_dedup_cls.return_value.consolidated_output_indices = set() + service._finalize_extracted_items([playbook], model_provenance=None) + + contexts = save.call_args.kwargs["lineage_contexts"] + assert len(contexts) == 1 + assert contexts[0] == LineageContext( + op_kind="create", + actor="extractor", + request_id="legacy-request", + model_name=None, + provider=None, + ) def test_run_manual_regular_no_window_size(mock_chat_completion): diff --git a/tests/server/services/playbook/test_playbook_generation_service_integration.py b/tests/server/services/playbook/test_playbook_generation_service_integration.py index 3051a1d99..0463df226 100644 --- a/tests/server/services/playbook/test_playbook_generation_service_integration.py +++ b/tests/server/services/playbook/test_playbook_generation_service_integration.py @@ -29,6 +29,7 @@ def disable_mock_llm_response(monkeypatch): PlaybookConfig, ) from reflexio.server.api_endpoints.request_context import RequestContext +from reflexio.server.llm._litellm_types import CompletionResult, ModelProvenance from reflexio.server.services.playbook.playbook_service_utils import ( PlaybookGenerationRequest, StructuredPlaybookContent, @@ -184,6 +185,11 @@ def mock_generate_chat_response(messages, **kwargs): service.client.generate_chat_response = MagicMock( side_effect=mock_generate_chat_response ) + service.client.generate_chat_response_with_provenance = MagicMock( + side_effect=lambda *args, **kwargs: CompletionResult( + mock_generate_chat_response(*args, **kwargs), ModelProvenance() + ) + ) @skip_in_precommit @@ -420,9 +426,20 @@ def mock_generate_chat_response(messages, **kwargs): ] ) - with patch( - "reflexio.server.llm.litellm_client.LiteLLMClient.generate_chat_response", - side_effect=mock_generate_chat_response, + def mock_generate_with_provenance(*args, **kwargs): + return CompletionResult( + mock_generate_chat_response(*args, **kwargs), ModelProvenance() + ) + + with ( + patch( + "reflexio.server.llm.litellm_client.LiteLLMClient.generate_chat_response", + side_effect=mock_generate_chat_response, + ), + patch( + "reflexio.server.llm.litellm_client.LiteLLMClient.generate_chat_response_with_provenance", + side_effect=mock_generate_with_provenance, + ), ): # Create playbook generation request with new API request = PlaybookGenerationRequest( diff --git a/tests/server/services/playbook_optimizer/test_optimizer_supersede_integration.py b/tests/server/services/playbook_optimizer/test_optimizer_supersede_integration.py index 6d6e18b59..f0b26bc59 100644 --- a/tests/server/services/playbook_optimizer/test_optimizer_supersede_integration.py +++ b/tests/server/services/playbook_optimizer/test_optimizer_supersede_integration.py @@ -94,10 +94,9 @@ def test_supersede_user_playbook_sets_superseded_by_and_revise_event(tmp_path): events = storage.get_lineage_events( entity_type="user_playbook", entity_id=str(result) ) - assert len(events) == 1 - assert events[0].op == "revise" - assert events[0].actor == "playbook_optimizer" - assert str(incumbent_id) in events[0].source_ids + assert [event.op for event in events] == ["create", "revise"] + assert events[1].actor == "playbook_optimizer" + assert str(incumbent_id) in events[1].source_ids def test_supersede_user_playbook_returns_none_for_non_current_incumbent(tmp_path): @@ -113,6 +112,7 @@ def test_supersede_user_playbook_returns_none_for_non_current_incumbent(tmp_path status=Status.ARCHIVED, # not CURRENT ) storage.save_user_playbooks([incumbent]) + events_before = storage.get_lineage_events(entity_type="user_playbook") playbooks_before = storage.conn.execute( "SELECT COUNT(*) as cnt FROM user_playbooks" @@ -134,9 +134,9 @@ def test_supersede_user_playbook_returns_none_for_non_current_incumbent(tmp_path ).fetchone()["cnt"] assert playbooks_after == playbooks_before, "no orphan row should remain" - # No lineage events should exist + # The failed successor contributes no row or event; the incumbent origin remains. events = storage.get_lineage_events(entity_type="user_playbook") - assert events == [] + assert events == events_before # --------------------------------------------------------------------------- @@ -189,10 +189,9 @@ def test_supersede_agent_playbook_sets_superseded_by_and_revise_event(tmp_path): events = storage.get_lineage_events( entity_type="agent_playbook", entity_id=str(result) ) - assert len(events) == 1 - assert events[0].op == "revise" - assert events[0].actor == "playbook_optimizer" - assert str(incumbent_id) in events[0].source_ids + assert [event.op for event in events] == ["create", "revise"] + assert events[1].actor == "playbook_optimizer" + assert str(incumbent_id) in events[1].source_ids def test_supersede_agent_playbook_returns_none_for_non_current_incumbent(tmp_path): @@ -210,6 +209,7 @@ def test_supersede_agent_playbook_returns_none_for_non_current_incumbent(tmp_pat ) ] ) + events_before = storage.get_lineage_events(entity_type="agent_playbook") agent_playbooks_before = storage.conn.execute( "SELECT COUNT(*) as cnt FROM agent_playbooks" @@ -233,7 +233,7 @@ def test_supersede_agent_playbook_returns_none_for_non_current_incumbent(tmp_pat ) events = storage.get_lineage_events(entity_type="agent_playbook") - assert events == [] + assert events == events_before # --------------------------------------------------------------------------- @@ -272,11 +272,10 @@ def test_supersede_user_playbook_revise_event_carries_job_request_id(tmp_path): events = storage.get_lineage_events( entity_type="user_playbook", entity_id=str(result) ) - assert len(events) == 1 - assert events[0].op == "revise" - assert events[0].request_id == run_id, ( + assert [event.op for event in events] == ["create", "revise"] + assert events[1].request_id == run_id, ( f"revise event must carry the job-derived run id {run_id!r}, " - f"got {events[0].request_id!r}" + f"got {events[1].request_id!r}" ) @@ -311,11 +310,10 @@ def test_supersede_agent_playbook_revise_event_carries_job_request_id(tmp_path): events = storage.get_lineage_events( entity_type="agent_playbook", entity_id=str(result) ) - assert len(events) == 1 - assert events[0].op == "revise" - assert events[0].request_id == run_id, ( + assert [event.op for event in events] == ["create", "revise"] + assert events[1].request_id == run_id, ( f"revise event must carry the job-derived run id {run_id!r}, " - f"got {events[0].request_id!r}" + f"got {events[1].request_id!r}" ) diff --git a/tests/server/services/profile/test_profile_consolidator.py b/tests/server/services/profile/test_profile_consolidator.py index e9fbab578..c1e614f4f 100644 --- a/tests/server/services/profile/test_profile_consolidator.py +++ b/tests/server/services/profile/test_profile_consolidator.py @@ -27,6 +27,7 @@ def disable_mock_llm_response(monkeypatch): ProfileTimeToLive, UserProfile, ) +from reflexio.server.llm._litellm_types import CompletionResult, ModelProvenance from reflexio.server.llm.litellm_client import ( LiteLLMClient, LiteLLMClientError, @@ -52,6 +53,16 @@ def disable_mock_llm_response(monkeypatch): def mock_llm_client(): """Create a mock LLM client.""" client = MagicMock(spec=LiteLLMClient) + + def _generate_with_provenance(*args, **kwargs): + value = client.generate_chat_response(*args, **kwargs) + if isinstance(value, CompletionResult): + return value + return CompletionResult(value=value, provenance=ModelProvenance()) + + client.generate_chat_response_with_provenance.side_effect = ( + _generate_with_provenance + ) client.get_embeddings.return_value = [[0.1] * 10, [0.2] * 10, [0.3] * 10] return client diff --git a/tests/server/services/profile/test_profile_generation_service.py b/tests/server/services/profile/test_profile_generation_service.py index 5c3f90f1b..4cdf2a140 100644 --- a/tests/server/services/profile/test_profile_generation_service.py +++ b/tests/server/services/profile/test_profile_generation_service.py @@ -22,6 +22,7 @@ ) from reflexio.models.config_schema import ProfileExtractorConfig from reflexio.server.api_endpoints.request_context import RequestContext +from reflexio.server.llm._litellm_types import ModelProvenance from reflexio.server.llm.litellm_client import LiteLLMClient, LiteLLMConfig from reflexio.server.services.base_generation_service import StatusChangeOperation from reflexio.server.services.profile.profile_generation_service_utils import ( @@ -310,9 +311,15 @@ def test_save_profiles(self, service, request_context, sample_profile): service._process_results([[sample_profile]]) - request_context.storage.add_user_profile.assert_called_once_with( - "user_1", [sample_profile], skip_embedding=True - ) + call = request_context.storage.add_user_profile.call_args + assert call.args[:2] == ("user_1", [sample_profile]) + assert call.kwargs["skip_embedding"] is True + contexts = call.kwargs["lineage_contexts"] + assert len(contexts) == 1 + assert contexts[0] is not None + assert contexts[0].op_kind == "create" + assert contexts[0].model_name is None + assert contexts[0].provider is None assert sample_profile.source == "api" assert sample_profile.status is None # CURRENT (not pending) @@ -330,6 +337,58 @@ def test_save_profiles_pending_status( assert sample_profile.status == Status.PENDING + def test_save_profiles_carries_extractor_model_provenance( + self, service, request_context, sample_profile + ): + self._setup_service_config(service) + service._last_model_provenance = ModelProvenance( + model_name="served-model", + provider="provider", + ) + + service._process_results([[sample_profile]]) + + context = request_context.storage.add_user_profile.call_args.kwargs[ + "lineage_contexts" + ][0] + assert context.op_kind == "create" + assert context.model_name == "served-model" + assert context.provider == "provider" + + def test_merged_profile_uses_consolidator_completion_provenance( + self, service, request_context, sample_profile + ): + self._setup_service_config(service) + merged = sample_profile.model_copy(update={"profile_id": "merged-profile"}) + provenance = ModelProvenance( + model_name="dedup-served", + provider="dedup-provider", + ) + + class FakeConsolidator: + model_provenance = provenance + lineage_sources_by_profile_id = {"merged-profile": ["old-profile"]} + consolidated_output_indices = {0} + + def __init__(self, **_kwargs): + pass + + def deduplicate(self, *_args): + return [merged], ["old-profile"], [] + + with patch( + "reflexio.server.services.profile.components.consolidator.ProfileConsolidator", + FakeConsolidator, + ): + service._process_results([[sample_profile]]) + + context = request_context.storage.add_user_profile.call_args.kwargs[ + "lineage_contexts" + ][0] + assert context.actor == "consolidator" + assert context.source_ids == ["old-profile"] + assert context.model_name == "dedup-served" + def test_save_failure_reraises_without_deleting( self, service, request_context, sample_profile ): @@ -361,9 +420,15 @@ def test_profiles_persisted_on_save_path( service._process_results([[sample_profile]]) - request_context.storage.add_user_profile.assert_called_once_with( - "user_1", [sample_profile], skip_embedding=True - ) + call = request_context.storage.add_user_profile.call_args + assert call.args[:2] == ("user_1", [sample_profile]) + assert call.kwargs["skip_embedding"] is True + contexts = call.kwargs["lineage_contexts"] + assert len(contexts) == 1 + assert contexts[0] is not None + assert contexts[0].op_kind == "create" + assert contexts[0].model_name is None + assert contexts[0].provider is None # =============================== diff --git a/tests/server/services/storage/sqlite_storage/test_create_lineage_provenance_integration.py b/tests/server/services/storage/sqlite_storage/test_create_lineage_provenance_integration.py new file mode 100644 index 000000000..c9831eb3a --- /dev/null +++ b/tests/server/services/storage/sqlite_storage/test_create_lineage_provenance_integration.py @@ -0,0 +1,187 @@ +from __future__ import annotations + +import time +from unittest.mock import patch + +import pytest + +import reflexio.server.services.storage.sqlite_storage.playbook._user as playbook_mod +import reflexio.server.services.storage.sqlite_storage.profiles._profile_store as profile_mod +from reflexio.models.api_schema.domain.entities import LineageContext +from reflexio.models.api_schema.service_schemas import UserPlaybook, UserProfile +from reflexio.server.services.storage.error import StorageError +from reflexio.server.services.storage.sqlite_storage import SQLiteStorage + +pytestmark = pytest.mark.integration + + +@pytest.fixture(autouse=True) +def _local_governance_secret(monkeypatch) -> None: + monkeypatch.setenv("REFLEXIO_GOVERNANCE_REF_SECRET", "test-governance-secret") + + +def _storage(tmp_path) -> SQLiteStorage: + storage = SQLiteStorage(org_id="org-create", db_path=str(tmp_path / "create.db")) + storage.migrate() + return storage + + +def _context() -> LineageContext: + return LineageContext( + op_kind="create", + actor="extractor", + request_id="req-create", + model_name="claude-sonnet-4-5-20250929", + provider="anthropic", + ) + + +def _profile() -> UserProfile: + return UserProfile( + profile_id="p-create", + user_id="u1", + content="likes concise answers", + last_modified_timestamp=int(time.time()), + generated_from_request_id="req-create", + ) + + +def _playbook() -> UserPlaybook: + return UserPlaybook( + user_id="u1", + agent_version="v1", + request_id="req-create", + content="Answer concisely.", + ) + + +def test_profile_create_event_is_atomic_and_provenance_aware(tmp_path) -> None: + storage = _storage(tmp_path) + storage.add_user_profile( + "u1", [_profile()], skip_embedding=True, lineage_contexts=[_context()] + ) + + event = storage.get_lineage_events(entity_type="profile", entity_id="p-create")[0] + assert event.op == "create" + assert event.actor == "extractor" + assert event.model_name == "claude-sonnet-4-5-20250929" + + +def test_profile_create_without_context_emits_lineage_with_unknown_model( + tmp_path, +) -> None: + storage = _storage(tmp_path) + storage.add_user_profile("u1", [_profile()], skip_embedding=True) + + event = storage.get_lineage_events(entity_type="profile", entity_id="p-create")[0] + assert event.op == "create" + assert event.request_id == "req-create" + assert event.model_name is None + assert event.provider is None + + +def test_profile_replace_does_not_emit_a_second_create(tmp_path) -> None: + storage = _storage(tmp_path) + profile = _profile() + storage.add_user_profile( + "u1", [profile], skip_embedding=True, lineage_contexts=[_context()] + ) + profile.content = "updated in place" + storage.add_user_profile( + "u1", [profile], skip_embedding=True, lineage_contexts=[_context()] + ) + + events = storage.get_lineage_events(entity_type="profile", entity_id="p-create") + assert [event.op for event in events] == ["create"] + + +def test_user_playbook_create_event_is_atomic_and_provenance_aware(tmp_path) -> None: + storage = _storage(tmp_path) + playbook = _playbook() + storage.save_user_playbooks( + [playbook], skip_embedding=True, lineage_contexts=[_context()] + ) + + event = storage.get_lineage_events( + entity_type="user_playbook", entity_id=str(playbook.user_playbook_id) + )[0] + assert event.op == "create" + assert event.provider == "anthropic" + + +def test_user_playbook_without_context_emits_lineage_with_unknown_model( + tmp_path, +) -> None: + storage = _storage(tmp_path) + playbook = _playbook() + + storage.save_user_playbooks([playbook], skip_embedding=True) + + event = storage.get_lineage_events( + entity_type="user_playbook", entity_id=str(playbook.user_playbook_id) + )[0] + assert event.op == "create" + assert event.request_id == "req-create" + assert event.model_name is None + assert event.provider is None + + +@pytest.mark.parametrize("kind", ["profile", "playbook"]) +def test_context_length_is_validated_before_db_work(tmp_path, kind: str) -> None: + storage = _storage(tmp_path) + with pytest.raises(StorageError, match="lineage_contexts must match"): + if kind == "profile": + storage.add_user_profile( + "u1", [_profile()], skip_embedding=True, lineage_contexts=[] + ) + else: + storage.save_user_playbooks( + [_playbook()], skip_embedding=True, lineage_contexts=[] + ) + assert storage.get_lineage_events(org_id="org-create") == [] + + +@pytest.mark.parametrize("kind", ["profile", "playbook"]) +def test_create_context_rejects_other_operation_kinds(tmp_path, kind: str) -> None: + storage = _storage(tmp_path) + context = _context().model_copy(update={"op_kind": "revise"}) + + with pytest.raises(StorageError, match="must use op_kind='create'"): + if kind == "profile": + storage.add_user_profile( + "u1", [_profile()], skip_embedding=True, lineage_contexts=[context] + ) + else: + storage.save_user_playbooks( + [_playbook()], skip_embedding=True, lineage_contexts=[context] + ) + assert storage.get_lineage_events(org_id="org-create") == [] + + +def test_profile_event_failure_rolls_back_insert(tmp_path) -> None: + storage = _storage(tmp_path) + with ( + patch.object( + profile_mod, "_append_event_stmt", side_effect=RuntimeError("boom") + ), + pytest.raises(StorageError, match="boom"), + ): + storage.add_user_profile( + "u1", [_profile()], skip_embedding=True, lineage_contexts=[_context()] + ) + assert storage.get_profile_by_id("p-create") is None + + +def test_playbook_event_failure_rolls_back_insert(tmp_path) -> None: + storage = _storage(tmp_path) + playbook = _playbook() + with ( + patch.object( + playbook_mod, "_append_event_stmt", side_effect=RuntimeError("boom") + ), + pytest.raises(StorageError, match="boom"), + ): + storage.save_user_playbooks( + [playbook], skip_embedding=True, lineage_contexts=[_context()] + ) + assert storage.get_user_playbooks(user_id="u1") == [] diff --git a/tests/server/services/storage/sqlite_storage/test_lineage_model_provenance_migration.py b/tests/server/services/storage/sqlite_storage/test_lineage_model_provenance_migration.py new file mode 100644 index 000000000..988c92422 --- /dev/null +++ b/tests/server/services/storage/sqlite_storage/test_lineage_model_provenance_migration.py @@ -0,0 +1,63 @@ +from __future__ import annotations + +import sqlite3 + +import pytest + +from reflexio.server.services.storage.sqlite_storage import SQLiteStorage + +pytestmark = pytest.mark.integration + +_MODEL_COLUMNS = { + "model_name", + "provider", +} + + +def test_fresh_database_has_model_provenance_columns(tmp_path) -> None: + storage = SQLiteStorage(org_id="org", db_path=str(tmp_path / "fresh.db")) + storage.migrate() + columns = { + row["name"] for row in storage.conn.execute("PRAGMA table_info(lineage_event)") + } + assert columns >= _MODEL_COLUMNS + assert "requested_model" not in columns + + +def test_legacy_database_upgrade_adds_nullable_columns_without_backfill( + tmp_path, +) -> None: + db_path = tmp_path / "legacy.db" + conn = sqlite3.connect(db_path) + conn.executescript(""" + CREATE TABLE lineage_event ( + event_id INTEGER PRIMARY KEY AUTOINCREMENT, + org_id TEXT NOT NULL, + entity_type TEXT NOT NULL, + entity_id TEXT NOT NULL, + op TEXT NOT NULL, + prov_relation TEXT NOT NULL DEFAULT '', + source_ids TEXT NOT NULL DEFAULT '[]', + actor TEXT NOT NULL DEFAULT '', + request_id TEXT NOT NULL DEFAULT '', + reason TEXT NOT NULL DEFAULT '', + created_at INTEGER NOT NULL, + UNIQUE (org_id, entity_type, entity_id, op, request_id) + ); + INSERT INTO lineage_event ( + org_id, entity_type, entity_id, op, created_at + ) VALUES ('org', 'profile', 'legacy-profile', 'create', 1); + """) + conn.commit() + conn.close() + + storage = SQLiteStorage(org_id="org", db_path=str(db_path)) + storage.migrate() + columns = { + row["name"] for row in storage.conn.execute("PRAGMA table_info(lineage_event)") + } + assert columns >= _MODEL_COLUMNS + + event = storage.get_lineage_events(entity_id="legacy-profile")[0] + assert event.model_name is None + assert event.provider is None diff --git a/tests/server/services/storage/sqlite_storage/test_playbook_atomicity_characterization_integration.py b/tests/server/services/storage/sqlite_storage/test_playbook_atomicity_characterization_integration.py index b9965ad6e..d9c68eb80 100644 --- a/tests/server/services/storage/sqlite_storage/test_playbook_atomicity_characterization_integration.py +++ b/tests/server/services/storage/sqlite_storage/test_playbook_atomicity_characterization_integration.py @@ -4,7 +4,7 @@ pin the CURRENT commit / lineage-event / no-op behavior of the top-risk, atomicity-sensitive methods so a "tidying" reorder during the mixin split is caught by a failing test. Modeled on the gold standard -``test_save_agent_playbook_with_aggregate_event_integration.py``. +``test_save_agent_playbooks_integration.py``. Methods characterized here (SQLite side): - ``supersede_user_playbooks_by_ids`` — soft-delete to SUPERSEDED, per-row diff --git a/tests/server/services/storage/sqlite_storage/test_save_agent_playbook_with_aggregate_event_integration.py b/tests/server/services/storage/sqlite_storage/test_save_agent_playbook_with_aggregate_event_integration.py index 08ff9c179..17bc13e62 100644 --- a/tests/server/services/storage/sqlite_storage/test_save_agent_playbook_with_aggregate_event_integration.py +++ b/tests/server/services/storage/sqlite_storage/test_save_agent_playbook_with_aggregate_event_integration.py @@ -1,4 +1,4 @@ -"""TDD tests for save_agent_playbook_with_aggregate_event — SQLite atomic write side. +"""SQLite atomic lineage tests for canonical agent-playbook saves. Five tests: 1. Happy path: row inserted + exactly one op=aggregate event with correct @@ -17,13 +17,27 @@ import pytest import reflexio.server.services.storage.sqlite_storage.playbook._agent as _agent_playbook_mod -from reflexio.models.api_schema.domain.entities import AgentPlaybook +from reflexio.models.api_schema.domain.entities import AgentPlaybook, LineageContext from reflexio.server.services.storage.error import StorageError from reflexio.server.services.storage.sqlite_storage import SQLiteStorage pytestmark = pytest.mark.integration +def _context( + *, source_ids: list[str] | None = None, request_id: str = "run-x" +) -> LineageContext: + return LineageContext( + op_kind="aggregate", + actor="aggregator", + source_ids=source_ids or [], + request_id=request_id, + reason="aggregate:full_archive", + model_name="claude-sonnet-4-5-20250929", + provider="anthropic", + ) + + # --------------------------------------------------------------------------- # Fixture # --------------------------------------------------------------------------- @@ -54,18 +68,15 @@ def _make_playbook( # --------------------------------------------------------------------------- -class TestSaveAgentPlaybookWithAggregateEvent: +class TestSaveAgentPlaybooksWithAggregateContext: def test_happy_path_row_and_event_both_written(self, tmp_path): """Row is inserted and exactly one aggregate event exists with correct fields.""" s = _store(tmp_path) pb = _make_playbook() - result = s.save_agent_playbook_with_aggregate_event( - pb, - source_ids=["10", "11"], - request_id="run-x", - run_mode="full_archive", - ) + result = s.save_agent_playbooks( + [pb], lineage_contexts=[_context(source_ids=["10", "11"])] + )[0] # Row exists and has a real ID assert result.agent_playbook_id > 0 @@ -86,6 +97,7 @@ def test_happy_path_row_and_event_both_written(self, tmp_path): assert ev.request_id == "run-x" assert ev.actor == "aggregator" assert ev.prov_relation == "wasDerivedFrom" + assert ev.model_name == "claude-sonnet-4-5-20250929" def test_atomicity_rollback_on_event_append_failure(self, tmp_path): """If _append_event_stmt raises, the INSERT is rolled back — no orphaned row.""" @@ -104,11 +116,9 @@ def test_atomicity_rollback_on_event_append_failure(self, tmp_path): # handle_exceptions wraps RuntimeError into StorageError pytest.raises(StorageError, match="simulated event failure"), ): - s.save_agent_playbook_with_aggregate_event( - pb, - source_ids=["1"], - request_id="run-fail", - run_mode="incremental", + s.save_agent_playbooks( + [pb], + lineage_contexts=[_context(source_ids=["1"], request_id="run-fail")], ) # Row count must be unchanged — INSERT rolled back @@ -127,18 +137,15 @@ def test_fts_indexes_new_playbook(self, tmp_path): s = _store(tmp_path) pb = _make_playbook(trigger="unique_trigger_xyz", content="some content") - result = s.save_agent_playbook_with_aggregate_event( - pb, - source_ids=[], - request_id="run-fts", - run_mode="incremental", - ) + result = s.save_agent_playbooks( + [pb], lineage_contexts=[_context(request_id="run-fts")] + )[0] hits = s.search_agent_playbooks( SearchAgentPlaybookRequest(query="unique_trigger_xyz", top_k=10) ) assert any(h.agent_playbook_id == result.agent_playbook_id for h in hits), ( - "New playbook not found in FTS index after save_agent_playbook_with_aggregate_event" + "New playbook not found in FTS index after canonical save" ) def test_index_failure_after_commit_does_not_rollback_row(self, tmp_path): @@ -156,12 +163,12 @@ def test_index_failure_after_commit_does_not_rollback_row(self, tmp_path): "_index_agent_playbook_fts_vec", side_effect=RuntimeError("simulated FTS index failure"), ): - result = s.save_agent_playbook_with_aggregate_event( - pb, - source_ids=["42"], - request_id="run-idx-fail", - run_mode="incremental", - ) + result = s.save_agent_playbooks( + [pb], + lineage_contexts=[ + _context(source_ids=["42"], request_id="run-idx-fail") + ], + )[0] # Method returns the saved playbook normally assert result is not None @@ -188,12 +195,10 @@ def test_empty_request_id_raises_before_write(self, tmp_path): rows_before = len(s.get_agent_playbooks()) # StorageError wraps ValueError via handle_exceptions - with pytest.raises((ValueError, StorageError), match="non-empty request_id"): - s.save_agent_playbook_with_aggregate_event( - pb, - source_ids=["1"], - request_id="", - run_mode="full_archive", + with pytest.raises((ValueError, StorageError), match="requires request_id"): + s.save_agent_playbooks( + [pb], + lineage_contexts=[_context(source_ids=["1"], request_id="")], ) rows_after = len(s.get_agent_playbooks()) diff --git a/tests/server/services/storage/test_lineage_b1_update_integration.py b/tests/server/services/storage/test_lineage_b1_update_integration.py index 2f9fd6d56..04627a32d 100644 --- a/tests/server/services/storage/test_lineage_b1_update_integration.py +++ b/tests/server/services/storage/test_lineage_b1_update_integration.py @@ -42,7 +42,7 @@ def test_update_user_playbook_content_emits_revise(tmp_path): ev = s.get_lineage_events( entity_id=str(pb.user_playbook_id), entity_type="user_playbook" ) - assert [e.op for e in ev] == ["revise"] + assert [e.op for e in ev] == ["create", "revise"] assert s.get_user_playbook_by_id(pb.user_playbook_id).content == "new guidance" @@ -54,7 +54,7 @@ def test_update_user_playbook_metadata_only_emits_status_change(tmp_path): ev = s.get_lineage_events( entity_id=str(pb.user_playbook_id), entity_type="user_playbook" ) - assert [e.op for e in ev] == ["status_change"] + assert [e.op for e in ev] == ["create", "status_change"] def test_update_user_playbook_multiple_edits_each_produce_event(tmp_path): @@ -67,7 +67,7 @@ def test_update_user_playbook_multiple_edits_each_produce_event(tmp_path): ev = s.get_lineage_events( entity_id=str(pb.user_playbook_id), entity_type="user_playbook" ) - assert [e.op for e in ev] == ["revise", "revise"] + assert [e.op for e in ev] == ["create", "revise", "revise"] def test_update_user_playbook_trigger_change_emits_revise(tmp_path): @@ -78,7 +78,7 @@ def test_update_user_playbook_trigger_change_emits_revise(tmp_path): ev = s.get_lineage_events( entity_id=str(pb.user_playbook_id), entity_type="user_playbook" ) - assert [e.op for e in ev] == ["revise"] + assert [e.op for e in ev] == ["create", "revise"] def test_read_user_playbook_as_of_for_learning_rejects_post_serve_revise(tmp_path): @@ -294,7 +294,7 @@ def test_update_agent_playbook_content_emits_revise(tmp_path): ev = s.get_lineage_events( entity_id=str(saved[0].agent_playbook_id), entity_type="agent_playbook" ) - assert [e.op for e in ev] == ["revise"] + assert [e.op for e in ev] == ["create", "revise"] def test_update_agent_playbook_metadata_only_emits_status_change(tmp_path): @@ -305,7 +305,7 @@ def test_update_agent_playbook_metadata_only_emits_status_change(tmp_path): ev = s.get_lineage_events( entity_id=str(saved[0].agent_playbook_id), entity_type="agent_playbook" ) - assert [e.op for e in ev] == ["status_change"] + assert [e.op for e in ev] == ["create", "status_change"] # --------------------------------------------------------------------------- @@ -321,7 +321,7 @@ def test_update_agent_playbook_status_always_emits_status_change(tmp_path): ev = s.get_lineage_events( entity_id=str(saved[0].agent_playbook_id), entity_type="agent_playbook" ) - assert [e.op for e in ev] == ["status_change"] + assert [e.op for e in ev] == ["create", "status_change"] # --------------------------------------------------------------------------- @@ -342,7 +342,7 @@ def test_update_user_profile_emits_revise(tmp_path): updated = profile.model_copy(update={"content": "updated content"}) s.update_user_profile_by_id("u", str(profile.profile_id), updated) ev = s.get_lineage_events(entity_id=str(profile.profile_id), entity_type="profile") - assert [e.op for e in ev] == ["revise"] + assert [e.op for e in ev] == ["create", "revise"] fetched = s.get_profile_by_id(str(profile.profile_id)) assert fetched is not None assert fetched.content == "updated content" @@ -377,11 +377,11 @@ def test_archive_agent_playbooks_by_ids_already_archived_no_event(tmp_path): # First archive emits one status_change event. s.archive_agent_playbooks_by_ids([apid]) first = s.get_lineage_events(entity_id=str(apid), entity_type="agent_playbook") - assert [e.op for e in first] == ["status_change"] + assert [e.op for e in first] == ["create", "status_change"] # Re-archiving the already-archived row must emit no further event. s.archive_agent_playbooks_by_ids([apid]) second = s.get_lineage_events(entity_id=str(apid), entity_type="agent_playbook") - assert [e.op for e in second] == ["status_change"] + assert [e.op for e in second] == ["create", "status_change"] def test_archive_agent_playbooks_by_playbook_name_already_archived_no_event(tmp_path): @@ -391,10 +391,10 @@ def test_archive_agent_playbooks_by_playbook_name_already_archived_no_event(tmp_ apid = saved[0].agent_playbook_id s.archive_agent_playbooks_by_playbook_name("arch-pb") first = s.get_lineage_events(entity_id=str(apid), entity_type="agent_playbook") - assert [e.op for e in first] == ["status_change"] + assert [e.op for e in first] == ["create", "status_change"] s.archive_agent_playbooks_by_playbook_name("arch-pb") second = s.get_lineage_events(entity_id=str(apid), entity_type="agent_playbook") - assert [e.op for e in second] == ["status_change"] + assert [e.op for e in second] == ["create", "status_change"] # --------------------------------------------------------------------------- diff --git a/tests/server/services/storage/test_playbook_base_aggregate_emit.py b/tests/server/services/storage/test_playbook_base_aggregate_emit.py deleted file mode 100644 index f63c6b989..000000000 --- a/tests/server/services/storage/test_playbook_base_aggregate_emit.py +++ /dev/null @@ -1,128 +0,0 @@ -"""Unit tests for AgentPlaybookStoreMixin.save_agent_playbook_with_aggregate_event base default. - -Tests the base-class default directly via an unbound-method call with a mock -self, so the SQLite override (which has its own tests) does not interfere. - -Two tests: - 1. Retry + loud: append_lineage_event always fails → retried - _AGGREGATE_EVENT_EMIT_ATTEMPTS times, capture_anomaly called with - level="error", method RETURNS the saved playbook (does not raise). - 2. Happy path: append succeeds on first call → called exactly once, - capture_anomaly NOT called, emitted event has correct op and reason. -""" - -from __future__ import annotations - -from unittest.mock import MagicMock, patch - -import pytest - -from reflexio.models.api_schema.domain.entities import AgentPlaybook -from reflexio.server.services.storage.storage_base.playbook._agent import ( - _AGGREGATE_EVENT_EMIT_ATTEMPTS, - AgentPlaybookStoreMixin, -) - - -def _make_saved_playbook() -> AgentPlaybook: - pb = AgentPlaybook( - playbook_name="test-pb", - agent_version="v2", - content="Do the thing.", - ) - pb.agent_playbook_id = 42 - return pb - - -def _make_mock_self(saved_pb: AgentPlaybook, append_side_effect=None) -> MagicMock: - """Build a minimal mock self that satisfies AgentPlaybookStoreMixin's attribute accesses.""" - mock_self = MagicMock() - mock_self.save_agent_playbooks.return_value = [saved_pb] - mock_self.org_id = "org-x" - if append_side_effect is not None: - mock_self.append_lineage_event.side_effect = append_side_effect - return mock_self - - -class TestPlaybookBaseAggregateEmit: - def test_retry_and_loud_on_persistent_failure(self): - """append fails every time → retried N times, capture_anomaly(level='error'), no raise.""" - saved_pb = _make_saved_playbook() - mock_self = _make_mock_self( - saved_pb, append_side_effect=RuntimeError("transient db error") - ) - - with patch( - "reflexio.server.services.storage.storage_base.playbook._agent.capture_anomaly" - ) as mock_capture: - result = AgentPlaybookStoreMixin.save_agent_playbook_with_aggregate_event( - mock_self, - AgentPlaybook(playbook_name="test-pb", agent_version="v2", content="x"), - source_ids=["1", "2"], - request_id="r-fail", - run_mode="full_archive", - ) - - # Method must return the saved playbook — never raise - assert result is saved_pb - - # append_lineage_event retried exactly _AGGREGATE_EVENT_EMIT_ATTEMPTS times - assert ( - mock_self.append_lineage_event.call_count == _AGGREGATE_EVENT_EMIT_ATTEMPTS - ) - - # capture_anomaly called once with level="error" - mock_capture.assert_called_once() - _, kwargs = mock_capture.call_args - assert kwargs.get("level") == "error" - - def test_happy_path_first_attempt_succeeds(self): - """append succeeds on first try → called once, capture_anomaly NOT called.""" - saved_pb = _make_saved_playbook() - mock_self = _make_mock_self(saved_pb) - - with patch( - "reflexio.server.services.storage.storage_base.playbook._agent.capture_anomaly" - ) as mock_capture: - result = AgentPlaybookStoreMixin.save_agent_playbook_with_aggregate_event( - mock_self, - AgentPlaybook(playbook_name="test-pb", agent_version="v2", content="x"), - source_ids=["10", "11"], - request_id="r-ok", - run_mode="full_archive", - ) - - assert result is saved_pb - - # append called exactly once — no retry on success - assert mock_self.append_lineage_event.call_count == 1 - - # Verify the emitted event has correct op and reason - (event,) = mock_self.append_lineage_event.call_args.args - assert event.op == "aggregate" - assert event.reason == "aggregate:full_archive" - assert event.prov_relation == "wasDerivedFrom" - assert event.actor == "aggregator" - assert event.source_ids == ["10", "11"] - assert event.request_id == "r-ok" - - # capture_anomaly must NOT be called on success - mock_capture.assert_not_called() - - def test_empty_request_id_raises_before_save(self): - """Empty request_id raises ValueError before any storage write (no orphan row).""" - saved_pb = _make_saved_playbook() - mock_self = _make_mock_self(saved_pb) - - with pytest.raises(ValueError, match="non-empty request_id"): - AgentPlaybookStoreMixin.save_agent_playbook_with_aggregate_event( - mock_self, - AgentPlaybook(playbook_name="test-pb", agent_version="v2", content="x"), - source_ids=["1"], - request_id="", - run_mode="full_archive", - ) - - # No storage call must have been made - mock_self.save_agent_playbooks.assert_not_called() - mock_self.append_lineage_event.assert_not_called() diff --git a/tests/server/services/storage/test_sqlite_lineage_event_integration.py b/tests/server/services/storage/test_sqlite_lineage_event_integration.py index 4a4bf4181..9d26566e2 100644 --- a/tests/server/services/storage/test_sqlite_lineage_event_integration.py +++ b/tests/server/services/storage/test_sqlite_lineage_event_integration.py @@ -31,6 +31,27 @@ def test_append_then_get(tmp_path): assert len(rows) == 1 and rows[0].op == "merge" and rows[0].created_at > 0 +def test_model_provenance_round_trips_and_unknown_stays_null(tmp_path): + s = SQLiteStorage(org_id="org-42", db_path=str(tmp_path / "t.db")) + s.migrate() + s.append_lineage_event( + _evt( + entity_id="with-model", + model_name="claude-sonnet-4-5-20250929", + provider="anthropic", + ) + ) + s.append_lineage_event(_evt(entity_id="unknown-model")) + + with_model = s.get_lineage_events(entity_id="with-model")[0] + assert with_model.model_name == "claude-sonnet-4-5-20250929" + assert with_model.provider == "anthropic" + + unknown = s.get_lineage_events(entity_id="unknown-model")[0] + assert unknown.model_name is None + assert unknown.provider is None + + def test_append_is_idempotent_on_unique_key(tmp_path): s = SQLiteStorage(org_id="org-42", db_path=str(tmp_path / "t.db")) s.migrate() diff --git a/tests/server/services/test_non_extraction_learning_metering.py b/tests/server/services/test_non_extraction_learning_metering.py index f963dbe47..87f9875de 100644 --- a/tests/server/services/test_non_extraction_learning_metering.py +++ b/tests/server/services/test_non_extraction_learning_metering.py @@ -182,7 +182,7 @@ def test_aggregation_records_attributed_learnings_generated() -> None: """Aggregation emits one entity-backed event per generated playbook. ``saved_playbook_list`` entries always carry a real ``agent_playbook_id`` - (``save_agent_playbook_with_aggregate_event`` raises rather than + (``save_agent_playbooks`` raises rather than returning a partial row) -- aggregator.py is the one caller with a clean, always-populated per-record id list, so it uses the entity-backed path (Task A3) rather than the count-only fallback. diff --git a/tests/server/services/test_profile_generation_service.py b/tests/server/services/test_profile_generation_service.py index f5a4684b1..46b3408da 100644 --- a/tests/server/services/test_profile_generation_service.py +++ b/tests/server/services/test_profile_generation_service.py @@ -29,6 +29,7 @@ def disable_mock_llm_response(monkeypatch): RerunProfileGenerationRequest, ) from reflexio.models.config_schema import ProfileExtractorConfig +from reflexio.server.llm._litellm_types import CompletionResult, ModelProvenance from reflexio.server.llm.litellm_client import LiteLLMClient, LiteLLMConfig from reflexio.server.services.generation_service import GenerationService from reflexio.server.services.profile.profile_generation_service_utils import ( @@ -82,10 +83,21 @@ def mock_generate_chat_response_side_effect(messages, **kwargs): # Fallback: non-structured JSON string (legacy non-loop callers). return '```json\n{\n "add": [{\n "content": "like sushi",\n "time_to_live": "one_month"\n }]\n}\n```' - # Mock the LLM client's generate_chat_response method - with patch( - "reflexio.server.llm.litellm_client.LiteLLMClient.generate_chat_response", - side_effect=mock_generate_chat_response_side_effect, + def mock_generate_with_provenance(*args, **kwargs): + return CompletionResult( + mock_generate_chat_response_side_effect(*args, **kwargs), + ModelProvenance(), + ) + + with ( + patch( + "reflexio.server.llm.litellm_client.LiteLLMClient.generate_chat_response", + side_effect=mock_generate_chat_response_side_effect, + ), + patch( + "reflexio.server.llm.litellm_client.LiteLLMClient.generate_chat_response_with_provenance", + side_effect=mock_generate_with_provenance, + ), ): yield @@ -328,20 +340,30 @@ def mock_generate_chat_response(messages, **kwargs): # This is the actual profile extraction call # Check if parse_structured_output is True in kwargs if kwargs.get("parse_structured_output", False): - # Return the parsed dict directly - return { - "add": [ - { - "content": "like Italian food and sushi", - "time_to_live": "one_month", - } + return StructuredProfilesOutput( + profiles=[ + ProfileAddItem( + content="like Italian food and sushi", + time_to_live="one_month", + ) ] - } + ) return '```json\n{\n "add": [{\n "content": "like Italian food and sushi",\n "time_to_live": "one_month"\n }]\n}\n```' - with patch( - "reflexio.server.llm.litellm_client.LiteLLMClient.generate_chat_response", - side_effect=mock_generate_chat_response, + def mock_generate_with_provenance(*args, **kwargs): + return CompletionResult( + mock_generate_chat_response(*args, **kwargs), ModelProvenance() + ) + + with ( + patch( + "reflexio.server.llm.litellm_client.LiteLLMClient.generate_chat_response", + side_effect=mock_generate_chat_response, + ), + patch( + "reflexio.server.llm.litellm_client.LiteLLMClient.generate_chat_response_with_provenance", + side_effect=mock_generate_with_provenance, + ), ): # Create profile generation request - extractors collect from storage profile_generation_request = ProfileGenerationRequest(