From 5a2cdd8dff177fadf1ef1a60af43a4ea901bf339 Mon Sep 17 00:00:00 2001 From: yilu331 <275352964+yilu331@users.noreply.github.com> Date: Sun, 26 Jul 2026 14:05:32 -0700 Subject: [PATCH] feat: make learning lineage origins model-aware Record observed model and provider metadata on lineage events while keeping unknown values null. Make canonical playbook and profile saves emit atomic origin events, carry typed completion provenance structurally with each generated entity, and roll back lost supersede races without ghost rows or events. --- reflexio/models/api_schema/domain/entities.py | 9 + .../server/llm/_litellm_text_generation.py | 181 +++++++- reflexio/server/llm/_litellm_types.py | 43 +- reflexio/server/llm/_provider_concurrency.py | 13 + .../llm/providers/claude_code_provider.py | 59 ++- .../providers/claude_code_stream_parser.py | 36 ++ reflexio/server/llm/tools.py | 25 +- .../base_generation/_extraction_lifecycle.py | 3 + .../services/base_generation_service.py | 3 + .../server/services/deferred_learning_plan.py | 5 + .../server/services/extraction/outcome.py | 4 + .../services/extraction/resumable_agent.py | 52 ++- .../services/extraction/resume_worker.py | 47 +- .../playbook/components/aggregator.py | 62 ++- .../playbook/components/consolidator.py | 18 +- .../services/playbook/components/extractor.py | 3 + .../services/playbook/playbook_edit_apply.py | 43 +- reflexio/server/services/playbook/service.py | 48 +- .../services/playbook_optimizer/optimizer.py | 59 +-- .../profile/components/consolidator.py | 22 +- .../services/profile/components/extractor.py | 3 + reflexio/server/services/profile/service.py | 45 +- .../services/storage/sqlite_storage/_base.py | 12 +- .../storage/sqlite_storage/_lineage.py | 17 +- .../storage/sqlite_storage/playbook/_agent.py | 148 +++--- .../storage/sqlite_storage/playbook/_user.py | 59 ++- .../sqlite_storage/profiles/_profile_store.py | 65 ++- .../services/storage/storage_base/__init__.py | 2 + .../storage/storage_base/playbook/_agent.py | 90 +--- .../storage/storage_base/playbook/_user.py | 6 +- .../storage_base/profiles/_profile_store.py | 6 +- .../test_contradiction_resolution_e2e.py | 7 +- .../consolidation/test_consolidation_eval.py | 11 +- tests/models/test_lineage_models.py | 31 +- tests/server/llm/test_claude_code_provider.py | 92 +++- .../llm/test_claude_code_stream_parser.py | 48 ++ .../server/llm/test_litellm_client_surface.py | 2 +- .../llm/test_litellm_client_tool_calls.py | 18 +- tests/server/llm/test_litellm_client_unit.py | 428 +++++++++++++++++- tests/server/llm/test_tools.py | 81 +++- .../llm/test_tools_multi_stage_integration.py | 7 +- .../test_compute_persist_split.py | 12 +- .../test_model_provenance_envelope.py | 90 ++++ .../extraction/test_resumable_agent.py | 16 +- .../test_aggregation_lineage_integration.py | 22 +- ...est_aggregation_soft_delete_integration.py | 6 +- .../test_apply_playbook_edit_integration.py | 4 +- .../playbook/test_cluster_change_detection.py | 83 ++-- .../test_consolidation_lineage_integration.py | 30 +- .../test_extractor_polarity_integration.py | 10 + .../playbook/test_playbook_aggregator.py | 171 ++++--- .../playbook/test_playbook_consolidator.py | 97 +++- .../test_playbook_consolidator_integration.py | 18 +- .../playbook/test_playbook_edit_apply.py | 18 +- .../test_playbook_generation_service.py | 64 ++- ...playbook_generation_service_integration.py | 23 +- .../test_optimizer_supersede_integration.py | 36 +- .../profile/test_profile_consolidator.py | 11 + .../test_profile_generation_service.py | 77 +++- ...t_create_lineage_provenance_integration.py | 187 ++++++++ ...test_lineage_model_provenance_migration.py | 63 +++ ..._atomicity_characterization_integration.py | 2 +- ...aybook_with_aggregate_event_integration.py | 71 +-- .../test_lineage_b1_update_integration.py | 24 +- .../test_playbook_base_aggregate_emit.py | 128 ------ .../test_sqlite_lineage_event_integration.py | 21 + .../test_non_extraction_learning_metering.py | 2 +- .../test_profile_generation_service.py | 52 ++- 68 files changed, 2537 insertions(+), 714 deletions(-) create mode 100644 tests/server/services/extraction/test_model_provenance_envelope.py create mode 100644 tests/server/services/storage/sqlite_storage/test_create_lineage_provenance_integration.py create mode 100644 tests/server/services/storage/sqlite_storage/test_lineage_model_provenance_migration.py delete mode 100644 tests/server/services/storage/test_playbook_base_aggregate_emit.py 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(