From c959ef0ccd4f6cb3bbb9fc9d0f9e6a7de75a4378 Mon Sep 17 00:00:00 2001 From: Evan Date: Thu, 8 Oct 2026 19:38:38 +0800 Subject: [PATCH 1/2] feat: add structured provider response contracts --- docs/architecture.md | 6 +- docs/index.md | 1 + docs/providers.md | 40 ++ .../plans/2026-10-08-v3-issue-228.md | 63 +++ openkyrozen/agent/runtime.py | 1 + openkyrozen/providers/__init__.py | 2 + openkyrozen/providers/anthropic.py | 10 +- openkyrozen/providers/base.py | 23 ++ openkyrozen/providers/bedrock.py | 10 +- openkyrozen/providers/calls.py | 50 ++- openkyrozen/providers/fallback.py | 16 +- openkyrozen/providers/google.py | 24 +- openkyrozen/providers/models.py | 182 ++++++++ openkyrozen/providers/ollama.py | 33 +- openkyrozen/providers/openai.py | 32 +- openkyrozen/providers/perplexity.py | 13 +- tests/test_provider_contracts.py | 388 ++++++++++++++++++ 17 files changed, 853 insertions(+), 41 deletions(-) create mode 100644 docs/superpowers/plans/2026-10-08-v3-issue-228.md create mode 100644 openkyrozen/providers/models.py create mode 100644 tests/test_provider_contracts.py diff --git a/docs/architecture.md b/docs/architecture.md index 3e8cc31..b91ebb9 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -74,7 +74,8 @@ requirements R236–R249. Each row names one package owner, not a new class to s The existing components are the migration starting points on V2; the linked issues implement the V3 behavior. In particular, native AgentEngine execution, separate Working/Strategic Plans, typed provider responses and typed runtime events are -future work, not capabilities delivered by this ownership change. +separate implementation work, not capabilities established by ownership names +alone. The structured provider foundation is described in [providers](providers.md). | Requirement / subsystem | Package owner and existing components | Responsibility and authority boundary | V3 implementation | | --- | --- | --- | --- | @@ -139,7 +140,8 @@ V3 architecture is already implemented. The inbound chat contract is `AgentRuntime.chat(session, message, *, clear_tasks=False, profile=None, memory_context=None, on_event=None, approve=None) -> str`. -Providers keep `LLMProvider.chat/chat_stream`. Small feature-owned protocols describe +Providers expose structured `LLMProvider.chat_response` and `get_capabilities`, +with checked `chat` tuple compatibility and existing text-only `chat_stream`. Small feature-owned protocols describe memory/vector, task, learning, scheduling, history and interaction storage; event and approval boundaries are callables. Concrete adapters are supplied at composition. Foreground, durable, MCP and subagent actions share the executor and produce the same diff --git a/docs/index.md b/docs/index.md index 378ba0f..b4ecd31 100644 --- a/docs/index.md +++ b/docs/index.md @@ -27,6 +27,7 @@ This index covers the current product guides, developer references, and dated en - [Reports index](reports.md) explains the dates, scope, and limits of recorded audits, validation runs, and benchmarks. - [Modular refactor validation](modular-refactor-validation.md), [production audit](production-audit-2026-10-04.md), [shipping audit](production-shipping-audit-2026-10-04.md), and [self-learning audit](self-learning-audit-2026-10-03.md) preserve historical evidence. - [V3 architecture ownership plan](superpowers/plans/2026-10-07-v3-issue-227.md) records the scoped implementation and verification for issue #227. +- [V3 provider contract plan](superpowers/plans/2026-10-08-v3-issue-228.md) records the structured response and compatibility migration for issue #228. - Machine-readable audit receipts and benchmark inputs remain beside their corresponding reports under `docs/` and `docs/benchmarks/`. The repository's [English overview](../README.md) links here. Technical guides are maintained in English; README translations provide localized entry points. diff --git a/docs/providers.md b/docs/providers.md index 825f3e8..6001e09 100644 --- a/docs/providers.md +++ b/docs/providers.md @@ -52,6 +52,46 @@ The normal provider setting has separate simple and complex model slots. The age The registry marks a subset of providers for automatic selection and records preferred fallback providers. A fallback is not guaranteed: it still needs credentials, an available compatible model, and a supported transport. Alternate sub-agent providers use their own configured environment credentials and do not inherit a generic main-provider key. See [sub-agents](subagents.md) before using cross-provider assignments. +## Structured response contract (V3 issue #228) + +`LLMProvider.chat_response(messages, model=None)` returns `ModelResponse`: +`text`, ordered `ToolCall` objects, normalized `finish_reason`, usage and explicit +provider/model/response metadata. `ToolCall` preserves a native call ID, name and +JSON-object arguments. Contract validation raises `ProviderContractError`, a +`ValueError` subclass; fallback does not replay a received malformed response. IDs omitted by a transport are generated within the response +scope and listed in `metadata["synthesized_call_ids"]`. Malformed arguments, +duplicate IDs/JSON keys and non-JSON values are rejected; calls are never rendered +as assistant text or executed by the provider boundary. + +Finish reasons distinguish `final`, `tool_request`, `length`, `provider_error`, +`cancelled`, `blocked` and `unknown`. The original status remains in metadata. +Absent or unfamiliar status does not prove completion. The existing `chat()` +tuple API delegates to this contract and refuses tool requests, incomplete, +blocked, cancelled or failed responses rather than silently losing their meaning. +Ordinary text retains its existing behavior. Direct Ollama `chat()` transport +failures now raise exceptions instead of returning an error-string tuple. Legacy-only providers are bridged +with unknown finish status and an explicit `legacy` marker. + +`get_capabilities(model=None)` reports flags for `native_tools`, `strict_schemas`, +`parallel_calls`, `text_streaming`, `streaming_tool_calls` and `reasoning_controls`. +These describe usable features of the shipped adapter API, not every feature an +underlying model might support. Text streaming reflects the existing implementation; +other flags remain false until the corresponding request paths are implemented. +Fallback capabilities are the conservative intersection of the configured, +model-mapped candidates. Structured responses retain the responding provider's +metadata. A successful response rejected by text conversion does not cause another +provider request. + +The runtime's `_get_model_response` preserves this object under the existing +bounded-call, usage and context scopes; `_get_llm_response` is the checked text +compatibility boundary. Existing stream methods remain text-only. Tool request +serialization, native history round trips and structured stream events belong to +[#229](https://github.com/EvanProgramming/OpenKyrozen/issues/229) and +[#230](https://github.com/EvanProgramming/OpenKyrozen/issues/230). This foundation +also relates to [provider adapter issue #225](https://github.com/EvanProgramming/OpenKyrozen/issues/225). +Offline fixtures and installed SDK types validate the contract; they do not establish +live provider interoperability. + ## Optional dependencies `pip install -e '.[web]'` installs FastAPI and Uvicorn. Provider extras are `.[claude]`, `.[gemini]`, `.[perplexity]`, `.[bedrock]`, and `.[vertex]`; `.[cloud]` groups the cloud integrations. `.[browser]` installs Playwright; browser execution also needs a browser installation. `.[all]` installs the supported optional Python integrations. The official installer has its own pinned set of dependencies; see [installation](installation.md). diff --git a/docs/superpowers/plans/2026-10-08-v3-issue-228.md b/docs/superpowers/plans/2026-10-08-v3-issue-228.md new file mode 100644 index 0000000..f054b5f --- /dev/null +++ b/docs/superpowers/plans/2026-10-08-v3-issue-228.md @@ -0,0 +1,63 @@ +# V3 issue #228: structured provider response contract + +Issue: https://github.com/EvanProgramming/OpenKyrozen/issues/228 (R001, R002, R004, R007) +Related: https://github.com/EvanProgramming/OpenKyrozen/issues/225 +Base: current main 25588796b4ed187dbde84b763042512c0e598405 +Branch: Evan/v3-228-provider-contracts + +## Task 1 — SDK-independent contract (RED → GREEN) + +Define ToolCall (nonempty id/name, JSON object arguments), ModelResponse (text, +ordered calls, usage, finish reason, explicit metadata), ProviderCapabilities and +FinishReason in providers/models.py. Normalize final/tool/length/error/cancelled/ +blocked/unknown outcomes. Never infer final from text or an absent finish status. +Preserve native IDs; synthesize response-scoped IDs only for transports omitting +IDs and identify them in metadata. Reject malformed/non-object/non-JSON arguments. +Tests must first fail for lost calls and finish reasons before adding code. + +## Task 2 — adapter and compatibility migration + +Add LLMProvider.chat_response() and get_capabilities(); keep chat() and text +streaming. Base chat_response bridges legacy-only providers with UNKNOWN status. +Native adapters expose full non-stream responses and chat delegates through the +shared checked compatibility conversion. That conversion refuses calls, tool +requests, length stops, errors, cancellation and blocked output. Preserve ordinary +text behavior and existing accounting exactly once, even for rejected responses. + +Implement each current transport: OpenAI chat/Azure, OpenAI Responses, Anthropic, +Google/Vertex, Bedrock, Ollama, Perplexity. Do not add tool request serialization. +Capabilities represent implemented adapter request features, not guessed endpoint +or model capabilities: native tools, strict schemas, parallel/streaming tools and +reasoning controls remain false until those paths exist; text streaming reflects +the actual override. Fallback uses model-mapped conservative intersections and +preserves the responding provider metadata; conversion failure is not a reason +to replay a successful structured response through another provider. + +## Task 3 — shared runtime boundary + +Add _get_model_response() to the existing providers/calls boundary and bind it +on AgentRuntime. Reuse bounded calls, scoped usage, context and token accounting. +_get_llm_response retains its string API through checked conversion. Legacy duck +providers remain supported, without treating dynamically invented mock attributes +as an explicit structured API. Current stream handling remains text-only; native +streaming is explicitly unsupported here and owned by #230. Do not invoke tools. + +## Task 4 — verification and delivery + +Cover all transports with text-only, mixed/tool-only, multiple-call and malformed +argument fixtures; validate IDs, usage, finish status, metadata and inheritance. +Use real installed SDK response types where available to check fixture shape. +Test legacy callers/providers, fallback identity/model mapping/capabilities, +lossy conversions, runtime structured receipt, no tool execution and cost-once. +Run focused tests on supported Python 3.12/3.13, full make test, make check, lint, +docs-check and offline agent/subagent acceptance. Obtain fresh independent review +and fix material findings. Sign/verify commits with GPG, create/attach one PR, +resolve valid review feedback and wait for final-head CI. Do not merge. + +## Limits + +No paid provider calls, new dependencies, raw SDK objects/credentials in metadata, +or changes to native request serialization (#229), registry (#232), structured +streaming (#230), or AgentEngine execution (#238). Live interoperability remains +unverified. Record this limitation in the PR rather than claiming fixture tests +are live acceptance. diff --git a/openkyrozen/agent/runtime.py b/openkyrozen/agent/runtime.py index 8c6daae..a5a927a 100644 --- a/openkyrozen/agent/runtime.py +++ b/openkyrozen/agent/runtime.py @@ -229,6 +229,7 @@ class AgentRuntime: _bounded_provider_call, _bounded_provider_stream, _get_llm_response, + _get_model_response, ) # agent.planner diff --git a/openkyrozen/providers/__init__.py b/openkyrozen/providers/__init__.py index 32a97e4..a968f72 100644 --- a/openkyrozen/providers/__init__.py +++ b/openkyrozen/providers/__init__.py @@ -14,3 +14,5 @@ from openkyrozen.providers.fallback import (FallbackProvider) from openkyrozen.providers.factory import (_PROVIDER_CLASSES, get_provider, get_fallback_provider, detect_provider, save_provider_config) from openkyrozen.security.credentials import (_get_encryption_key, _get_fernet, encrypt_api_key, decrypt_api_key, save_provider_config_encrypted) + +from openkyrozen.providers.models import ModelResponse, ToolCall, ProviderCapabilities, FinishReason, ProviderContractError diff --git a/openkyrozen/providers/anthropic.py b/openkyrozen/providers/anthropic.py index 1078a38..21ffb65 100644 --- a/openkyrozen/providers/anthropic.py +++ b/openkyrozen/providers/anthropic.py @@ -5,6 +5,7 @@ import time from typing import Any, Iterator from openkyrozen.providers.base import LLMProvider +from openkyrozen.providers.models import ModelResponse, model_response from openkyrozen.providers.config import ProviderConfig from openkyrozen.providers.retry import _retry_with_backoff @@ -39,6 +40,9 @@ def _prepare_messages(self, messages): return system_prompts, claude_messages def chat(self, messages: list[dict[str, str]], model: str | None = None) -> tuple[str, dict | None]: + return self.chat_response(messages, model).as_legacy_tuple() + + def chat_response(self, messages: list[dict[str, str]], model: str | None = None) -> ModelResponse: model = model or self.config.model_simple system_prompts, claude_messages = self._prepare_messages(messages) started = time.monotonic() @@ -68,7 +72,11 @@ def _call(): } usage_ledger._track_cost(self.config.provider, usage_dict, model=getattr(response, "model", None) or model, latency_ms=round((time.monotonic() - started) * 1000)) - return text.strip(), usage_dict + calls = [(getattr(block, "id", None), getattr(block, "name", None), getattr(block, "input", None)) + for block in response.content if getattr(block, "type", None) == "tool_use"] + return model_response(provider=self.name, model=model, actual_model=getattr(response, "model", None), text=text, usage=usage_dict, + calls=calls, response_id=getattr(response, "id", None), + raw_finish_reason=getattr(response, "stop_reason", None)) def chat_stream(self, messages: list[dict[str, str]], model: str | None = None) -> Iterator[str]: model = model or self.config.model_simple diff --git a/openkyrozen/providers/base.py b/openkyrozen/providers/base.py index 28494e9..6d1e649 100644 --- a/openkyrozen/providers/base.py +++ b/openkyrozen/providers/base.py @@ -2,6 +2,8 @@ from typing import Iterator from abc import ABC, abstractmethod +from inspect import getattr_static +from openkyrozen.providers.models import ModelResponse, ProviderCapabilities, ProviderContractError from openkyrozen.providers.config import ProviderConfig class LLMProvider(ABC): @@ -15,6 +17,16 @@ def chat(self, messages: list[dict[str, str]], model: str | None = None) -> tupl """Send messages to the LLM. Returns (content, usage_dict_or_None).""" ... + def chat_response(self, messages: list[dict], model: str | None = None) -> ModelResponse: + """Bridge providers implementing only the existing text contract.""" + text, usage = self.chat(messages, model) + return ModelResponse(text=text, usage=usage, + metadata={"provider": self.name, "model": model or self.config.model_simple, "legacy": True}) + + def get_capabilities(self, model: str | None = None) -> ProviderCapabilities: + """Only advertise features usable through the shipped adapter API.""" + return ProviderCapabilities(text_streaming=type(self).chat_stream is not LLMProvider.chat_stream) + def chat_stream(self, messages: list[dict[str, str]], model: str | None = None) -> Iterator[str]: """Stream response tokens. Default: fall back to non-streaming chat().""" text, _ = self.chat(messages, model) @@ -23,3 +35,14 @@ def chat_stream(self, messages: list[dict[str, str]], model: str | None = None) @property def name(self) -> str: return self.config.provider + + +def get_model_response(provider, messages, model=None) -> ModelResponse: + """Use an explicitly supplied structured API, or bridge a legacy duck provider.""" + if callable(getattr_static(provider, "chat_response", None)): + response = provider.chat_response(messages, model) + if not isinstance(response, ModelResponse): + raise ProviderContractError("chat_response must return ModelResponse") + return response + text, usage = provider.chat(messages, model) + return ModelResponse(text=text, usage=usage, metadata={"legacy": True}) diff --git a/openkyrozen/providers/bedrock.py b/openkyrozen/providers/bedrock.py index 3d7e07f..eb20115 100644 --- a/openkyrozen/providers/bedrock.py +++ b/openkyrozen/providers/bedrock.py @@ -6,6 +6,7 @@ import time from typing import Any, Iterator from openkyrozen.providers.base import LLMProvider +from openkyrozen.providers.models import ModelResponse, model_response from openkyrozen.providers.config import ProviderConfig from openkyrozen.providers.retry import _retry_with_backoff @@ -49,6 +50,9 @@ def _usage(data: dict[str, Any] | None) -> dict[str, int | None] | None: } def chat(self, messages: list[dict[str, str]], model: str | None = None) -> tuple[str, dict | None]: + return self.chat_response(messages, model).as_legacy_tuple() + + def chat_response(self, messages: list[dict[str, str]], model: str | None = None) -> ModelResponse: model = model or self.config.model_simple conversation, system = self._request(messages) started = time.monotonic() @@ -65,7 +69,11 @@ def chat(self, messages: list[dict[str, str]], model: str | None = None) -> tupl usage = self._usage(response) usage_ledger._track_cost(self.config.provider, usage, model=model, latency_ms=round((time.monotonic() - started) * 1000)) - return text.strip(), usage + calls = [(item["toolUse"].get("toolUseId"), item["toolUse"].get("name"), item["toolUse"].get("input")) + for item in content if isinstance(item, dict) and "toolUse" in item] + return model_response(provider=self.name, model=model, text=text, usage=usage, + calls=calls, response_id=response.get("ResponseMetadata", {}).get("RequestId"), + raw_finish_reason=response.get("stopReason")) def chat_stream(self, messages: list[dict[str, str]], model: str | None = None) -> Iterator[str]: model = model or self.config.model_simple diff --git a/openkyrozen/providers/calls.py b/openkyrozen/providers/calls.py index eb4a7e3..2b67637 100644 --- a/openkyrozen/providers/calls.py +++ b/openkyrozen/providers/calls.py @@ -8,6 +8,8 @@ import time from typing import Any from openkyrozen.providers.usage import usage_scope +from openkyrozen.providers.base import get_model_response +from openkyrozen.providers.models import ModelResponse from openkyrozen.agent.types import ContextOverflowError, ProviderUnavailableError @@ -102,6 +104,37 @@ def consume() -> None: raise TimeoutError(f"Provider timed out after {timeout:g}s") from exc +def _get_model_response(self, messages: list[dict], model: str | None = None) -> ModelResponse: + """Receive the full response under existing deadline, usage and context scopes.""" + provider = self.execution_context.provider or self.llm_provider + if provider is None: + raise ProviderUnavailableError(self.PROVIDER_UNAVAILABLE_MESSAGE) + self._last_prompt_tokens = self._last_completion_tokens = 0 + with usage_scope(store=self.memory_bank.store, user_id=self.memory_bank.user_id, + workspace_id=self.memory_bank.workspace_id, session_id=self.memory_bank.session_id, + run_id=self._active_usage_run_id.get(), surface=self._EXECUTION_SURFACE): + response = self._bounded_provider_call( + lambda: get_model_response(provider, messages, model or self.DEEPSEEK_MODEL) + ) + usage = response.usage or {} + self._last_prompt_tokens = usage.get("prompt_tokens", 0) or 0 + self._last_completion_tokens = usage.get("completion_tokens", 0) or 0 + if not self.execution_context.child_run_id: + self._total_prompt_tokens += self._last_prompt_tokens + self._total_completion_tokens += self._last_completion_tokens + if response.text.startswith("[Ollama Error]") and self._is_context_overflow_error(RuntimeError(response.text)): + raise RuntimeError(response.text) + state = self._active_context_state.get() + if state is not None and not self._in_context_compaction.get(): + if usage and not usage.get("_estimated"): + state.note_provider_usage(messages, int(self._last_prompt_tokens)) + self._cache_reported_context_tokens(messages, int(self._last_prompt_tokens)) + else: + state.update(messages) + self._store_context_status(state) + return response + + def _get_llm_response(self, messages: list[dict[str, str]], model: str | None = None, stream: bool = False, on_chunk: Any = None, on_stream_end: Any = None) -> str: provider = self.execution_context.provider or self.llm_provider @@ -137,22 +170,7 @@ def _get_llm_response(self, messages: list[dict[str, str]], model: str | None = if on_stream_end: on_stream_end() else: - text, usage_dict = self._bounded_provider_call( - lambda: provider.chat(messages, model or self.DEEPSEEK_MODEL) - ) - if (isinstance(text, str) and text.startswith("[Ollama Error]") - and self._is_context_overflow_error(RuntimeError(text))): - raise RuntimeError(text) - if usage_dict: - self._last_prompt_tokens = usage_dict.get("prompt_tokens", 0) - self._last_completion_tokens = usage_dict.get("completion_tokens", 0) - if not self.execution_context.child_run_id: - self._total_prompt_tokens += self._last_prompt_tokens - self._total_completion_tokens += self._last_completion_tokens - provider_reported_usage = not bool(usage_dict.get("_estimated")) - else: - self._last_prompt_tokens = 0 - self._last_completion_tokens = 0 + return self._get_model_response(messages, model).as_legacy_tuple()[0] except TimeoutError as exc: return f"[LLM Error] {exc}" except Exception as exc: diff --git a/openkyrozen/providers/fallback.py b/openkyrozen/providers/fallback.py index 62d69de..b459af1 100644 --- a/openkyrozen/providers/fallback.py +++ b/openkyrozen/providers/fallback.py @@ -2,7 +2,8 @@ import os from typing import Iterator -from openkyrozen.providers.base import LLMProvider +from openkyrozen.providers.base import LLMProvider, get_model_response +from openkyrozen.providers.models import ModelResponse, ProviderCapabilities, ProviderContractError from openkyrozen.providers.registry import PROVIDER_DEFAULT_MODELS, PROVIDER_ENV_VARS, PROVIDER_FALLBACKS from openkyrozen.providers.config import ProviderConfig, provider_is_configured from openkyrozen.providers.factory import get_provider @@ -79,12 +80,23 @@ def _raise_all_failed(attempts: list[tuple[str, Exception]]) -> None: error = RuntimeError(f"All providers failed ({details})") raise error from attempts[0][1] + def get_capabilities(self, model: str | None = None) -> ProviderCapabilities: + return ProviderCapabilities.intersection( + provider.get_capabilities(self._model_for(provider, model)) + for provider in [self._primary] + self._fallbacks + ) + def chat(self, messages: list[dict[str, str]], model: str | None = None) -> tuple[str, dict | None]: + return self.chat_response(messages, model).as_legacy_tuple() + + def chat_response(self, messages: list[dict[str, str]], model: str | None = None) -> ModelResponse: providers = [self._primary] + self._fallbacks attempts: list[tuple[str, Exception]] = [] for prov in providers: try: - return prov.chat(messages, self._model_for(prov, model)) + return get_model_response(prov, messages, self._model_for(prov, model)) + except ProviderContractError: + raise except Exception as e: attempts.append((prov.name, e)) self._raise_all_failed(attempts) diff --git a/openkyrozen/providers/google.py b/openkyrozen/providers/google.py index 63ed5e5..e53db82 100644 --- a/openkyrozen/providers/google.py +++ b/openkyrozen/providers/google.py @@ -6,6 +6,7 @@ import time from typing import Any, Iterator from openkyrozen.providers.base import LLMProvider +from openkyrozen.providers.models import ModelResponse, model_response from openkyrozen.providers.config import ProviderConfig from openkyrozen.providers.retry import _retry_with_backoff @@ -53,6 +54,9 @@ def _usage(response: Any) -> dict[str, int | None] | None: } def chat(self, messages: list[dict[str, str]], model: str | None = None) -> tuple[str, dict | None]: + return self.chat_response(messages, model).as_legacy_tuple() + + def chat_response(self, messages: list[dict[str, str]], model: str | None = None) -> ModelResponse: model = model or self.config.model_simple started = time.monotonic() @@ -67,11 +71,27 @@ def _call(): ) response = _retry_with_backoff(_call) - text = str(getattr(response, "text", "") or "") + candidates = getattr(response, "candidates", None) or () + candidate = candidates[0] if candidates else None + parts = getattr(getattr(candidate, "content", None), "parts", None) + text = ("".join(part.text for part in parts if getattr(part, "text", None) + and not getattr(part, "thought", False)) if parts is not None + else str(getattr(response, "text", "") or "")) usage_dict = self._usage(response) usage_ledger._track_cost(self.config.provider, usage_dict, model=model, latency_ms=round((time.monotonic() - started) * 1000)) - return text.strip(), usage_dict + calls = [] + for part in parts or (): + call = getattr(part, "function_call", None) + if call is not None: + arguments = getattr(call, "args", None) + calls.append((getattr(call, "id", None), getattr(call, "name", None), + {} if arguments is None else arguments)) + block_reason = getattr(getattr(response, "prompt_feedback", None), "block_reason", None) + return model_response(provider=self.name, model=model, actual_model=getattr(response, "model_version", None), text=text, usage=usage_dict, + calls=calls, response_id=getattr(response, "response_id", None), + raw_finish_reason=getattr(candidate, "finish_reason", None) or block_reason, + blocked=block_reason not in {None, "BLOCK_REASON_UNSPECIFIED"}) def chat_stream(self, messages: list[dict[str, str]], model: str | None = None) -> Iterator[str]: model = model or self.config.model_simple diff --git a/openkyrozen/providers/models.py b/openkyrozen/providers/models.py new file mode 100644 index 0000000..f7aa4bd --- /dev/null +++ b/openkyrozen/providers/models.py @@ -0,0 +1,182 @@ +"""Provider-independent non-streaming contracts; no SDK imports or execution.""" +from __future__ import annotations + +import copy +import json +import math +import uuid +from dataclasses import dataclass, field, fields +from enum import StrEnum +from typing import Any + + +class ProviderContractError(ValueError): + """A received response cannot satisfy the contract; do not replay generation.""" + + +class FinishReason(StrEnum): + FINAL = "final" + TOOL_REQUEST = "tool_request" + LENGTH = "length" + ERROR = "provider_error" + CANCELLED = "cancelled" + BLOCKED = "blocked" + UNKNOWN = "unknown" + + +def _validate_json(value: Any) -> None: + if isinstance(value, dict): + if any(not isinstance(key, str) for key in value): + raise ProviderContractError("Tool argument keys must be strings") + for item in value.values(): + _validate_json(item) + elif isinstance(value, list): + for item in value: + _validate_json(item) + elif value is None or isinstance(value, (str, bool, int)): + return + elif not isinstance(value, float) or not math.isfinite(value): + raise ProviderContractError("Tool arguments must contain finite JSON values") + + +def _unique_object(pairs): + result = {} + for key, value in pairs: + if key in result: + raise ProviderContractError(f"Duplicate tool argument key: {key}") + result[key] = value + return result + + +@dataclass(frozen=True) +class ToolCall: + id: str + name: str + arguments: dict[str, Any] + + def __post_init__(self): + if not isinstance(self.id, str) or not self.id.strip(): + raise ProviderContractError("Tool call id must be a nonempty string") + if not isinstance(self.name, str) or not self.name.strip(): + raise ProviderContractError("Tool call name must be a nonempty string") + if not isinstance(self.arguments, dict): + raise ProviderContractError("Tool arguments must be a JSON object") + try: + _validate_json(self.arguments) + object.__setattr__(self, "arguments", copy.deepcopy(self.arguments)) + except RecursionError as exc: + raise ProviderContractError("Tool arguments exceed supported JSON nesting") from exc + + +@dataclass(frozen=True) +class ProviderCapabilities: + native_tools: bool = False + strict_schemas: bool = False + parallel_calls: bool = False + text_streaming: bool = False + streaming_tool_calls: bool = False + reasoning_controls: bool = False + + def __post_init__(self): + if any(type(getattr(self, item.name)) is not bool for item in fields(self)): + raise ProviderContractError("Provider capability flags must be booleans") + + @classmethod + def intersection(cls, capabilities): + capabilities = tuple(capabilities) + return cls(**{item.name: bool(capabilities) and all(getattr(value, item.name) for value in capabilities) + for item in fields(cls)}) + + +@dataclass(frozen=True) +class ModelResponse: + text: str = "" + tool_calls: tuple[ToolCall, ...] = () + usage: dict[str, Any] | None = None + finish_reason: FinishReason = FinishReason.UNKNOWN + metadata: dict[str, Any] = field(default_factory=dict) + + def __post_init__(self): + if not isinstance(self.text, str) or not isinstance(self.finish_reason, FinishReason): + raise ProviderContractError("Invalid model response text or finish reason") + if not isinstance(self.tool_calls, (tuple, list)): + raise ProviderContractError("Model response calls must be an ordered collection") + object.__setattr__(self, "tool_calls", tuple(self.tool_calls)) + if any(not isinstance(call, ToolCall) for call in self.tool_calls): + raise ProviderContractError("Model response calls must be ToolCall objects") + if self.finish_reason == FinishReason.TOOL_REQUEST and not self.tool_calls: + raise ProviderContractError("Tool request has no calls") + if len({call.id for call in self.tool_calls}) != len(self.tool_calls): + raise ProviderContractError("Duplicate tool call ids") + if self.usage is not None and not isinstance(self.usage, dict): + raise ProviderContractError("Model response usage must be an object or None") + if not isinstance(self.metadata, dict): + raise ProviderContractError("Model response metadata must be an object") + try: + _validate_json(self.metadata) + except RecursionError as exc: + raise ProviderContractError("Metadata exceeds supported JSON nesting") from exc + + def as_legacy_tuple(self) -> tuple[str, dict | None]: + """Refuse to hide calls or unsuccessful completion in a text-only API.""" + if (self.tool_calls or self.finish_reason not in {FinishReason.FINAL, FinishReason.UNKNOWN} + or self.metadata.get("raw_finish_reason") in ("incomplete", "in_progress", "queued", "pause_turn")): + raise ProviderContractError(f"Structured response ({self.finish_reason}) requires chat_response(); text conversion refused") + return (self.text if self.metadata.get("legacy") else self.text.strip()), self.usage + + +def normalize_finish(reason, *, has_calls=False, detail=None) -> FinishReason: + reason = getattr(reason, "value", reason) + reason = reason.lower() if isinstance(reason, str) else "" + if reason == "incomplete": + reason = detail.lower() if isinstance(detail, str) else reason + if reason in {"stop", "completed", "end_turn", "stop_sequence"}: + return FinishReason.TOOL_REQUEST if has_calls else FinishReason.FINAL + if reason in {"tool_calls", "function_call", "tool_use"}: + return FinishReason.TOOL_REQUEST + if reason in {"length", "max_tokens", "max_output_tokens", "max_token", "model_context_window_exceeded"}: + return FinishReason.LENGTH + if reason in {"failed", "error", "malformed_function_call", "unexpected_tool_call"}: + return FinishReason.ERROR + if reason in {"cancelled", "canceled"}: + return FinishReason.CANCELLED + if reason in {"content_filter", "content_filtered", "guardrail_intervened", "safety", "recitation", + "refusal", "blocked", "blocklist", "prohibited_content", "spii"}: + return FinishReason.BLOCKED + return FinishReason.UNKNOWN + + +def model_response(*, provider, model, text, usage, raw_finish_reason=None, + calls=(), response_id=None, finish_detail=None, blocked=False, actual_model=None) -> ModelResponse: + """Build canonical calls from adapter-extracted (id, name, arguments) triples.""" + scope = response_id if isinstance(response_id, str) and response_id else "local_" + uuid.uuid4().hex + parsed, synthesized = [], [] + for index, (call_id, name, arguments) in enumerate(calls): + if call_id is None: + call_id = f"call_{scope}_{index}" + synthesized.append(call_id) + if isinstance(arguments, str): + try: + arguments = json.loads(arguments, object_pairs_hook=_unique_object) + except (ValueError, RecursionError) as exc: + raise ProviderContractError("Invalid serialized tool arguments") from exc + parsed.append(ToolCall(call_id, name, arguments)) + raw = getattr(raw_finish_reason, "value", raw_finish_reason) + metadata = {"provider": provider, "model": actual_model if isinstance(actual_model, str) and actual_model else model, "response_id": scope, + "raw_finish_reason": raw if isinstance(raw, str) else None} + if isinstance(finish_detail, str): + metadata["finish_detail"] = finish_detail + if synthesized: + metadata["synthesized_call_ids"] = synthesized + return ModelResponse(text, tuple(parsed), usage, + FinishReason.BLOCKED if blocked else normalize_finish(raw, has_calls=bool(parsed), detail=finish_detail), metadata) + + +def responses_output(response): + """Extract Responses API function calls without using its flattened text.""" + calls = [] + for item in getattr(response, "output", None) or (): + if getattr(item, "type", None) == "function_call": + calls.append((getattr(item, "call_id", None), getattr(item, "name", None), + getattr(item, "arguments", None))) + return calls diff --git a/openkyrozen/providers/ollama.py b/openkyrozen/providers/ollama.py index 4be352c..c5ef254 100644 --- a/openkyrozen/providers/ollama.py +++ b/openkyrozen/providers/ollama.py @@ -4,6 +4,7 @@ import sys import time from openkyrozen.providers.base import LLMProvider +from openkyrozen.providers.models import ModelResponse, model_response from openkyrozen.providers.config import ProviderConfig @@ -20,21 +21,25 @@ def __init__(self, config: ProviderConfig) -> None: self._base = config.base_url.replace("/v1", "") or "http://localhost:11434" def chat(self, messages: list[dict[str, str]], model: str | None = None) -> tuple[str, dict | None]: + return self.chat_response(messages, model).as_legacy_tuple() + + def chat_response(self, messages: list[dict[str, str]], model: str | None = None) -> ModelResponse: model = model or self.config.model_simple url = f"{self._base}/api/chat" payload = {"model": model, "messages": messages, "stream": False, "think": False} started = time.monotonic() - try: - resp = self._requests.post(url, json=payload, timeout=120) - resp.raise_for_status() - data = resp.json() - text = data.get("message", {}).get("content", "") - usage_dict = { - "prompt_tokens": data.get("prompt_eval_count", 0) or 0, - "completion_tokens": data.get("eval_count", 0) or 0, - } - usage_ledger._track_cost("ollama", usage_dict, model=model, - latency_ms=round((time.monotonic() - started) * 1000)) - return text.strip(), usage_dict - except Exception as e: - return f"[Ollama Error] {e}", None + resp = self._requests.post(url, json=payload, timeout=120) + resp.raise_for_status() + data = resp.json() + text = data.get("message", {}).get("content", "") + usage_dict = { + "prompt_tokens": data.get("prompt_eval_count", 0) or 0, + "completion_tokens": data.get("eval_count", 0) or 0, + } + usage_ledger._track_cost("ollama", usage_dict, model=model, + latency_ms=round((time.monotonic() - started) * 1000)) + calls = [(call.get("id"), call.get("function", {}).get("name"), + call.get("function", {}).get("arguments", {})) + for call in data.get("message", {}).get("tool_calls", [])] + return model_response(provider=self.name, model=model, text=text, usage=usage_dict, + calls=calls, raw_finish_reason="error" if data.get("error") else data.get("done_reason")) diff --git a/openkyrozen/providers/openai.py b/openkyrozen/providers/openai.py index c913609..62ba343 100644 --- a/openkyrozen/providers/openai.py +++ b/openkyrozen/providers/openai.py @@ -5,6 +5,7 @@ import time from typing import Any, Iterator from openkyrozen.providers.base import LLMProvider +from openkyrozen.providers.models import ModelResponse, model_response, responses_output, ProviderContractError from openkyrozen.providers.config import ProviderConfig, _provider_env_key from openkyrozen.providers.usage import _openai_usage_dict from openkyrozen.providers.retry import _retry_with_backoff @@ -29,6 +30,9 @@ def __init__(self, config: ProviderConfig) -> None: self._client = OpenAI(**kwargs) def chat(self, messages: list[dict[str, str]], model: str | None = None) -> tuple[str, dict | None]: + return self.chat_response(messages, model).as_legacy_tuple() + + def chat_response(self, messages: list[dict[str, str]], model: str | None = None) -> ModelResponse: model = model or self.config.model_simple started = time.monotonic() @@ -44,7 +48,21 @@ def _call(): usage_dict = _openai_usage_dict(usage) usage_ledger._track_cost(self.config.provider, usage_dict, model=getattr(response, "model", None) or model, latency_ms=round((time.monotonic() - started) * 1000)) - return text.strip(), usage_dict + message = response.choices[0].message + calls = [] + for call in getattr(message, "tool_calls", None) or (): + function = getattr(call, "function", None) + if function is None: + raise ProviderContractError("Unsupported native tool call type") + calls.append((getattr(call, "id", None), getattr(function, "name", None), + getattr(function, "arguments", None))) + legacy_call = getattr(message, "function_call", None) + if legacy_call is not None and not calls: + calls.append((None, getattr(legacy_call, "name", None), getattr(legacy_call, "arguments", None))) + return model_response(provider=self.name, model=model, actual_model=getattr(response, "model", None), text=text, usage=usage_dict, + calls=calls, response_id=getattr(response, "id", None), + raw_finish_reason=getattr(response.choices[0], "finish_reason", None), + blocked=bool(getattr(message, "refusal", None))) def chat_stream(self, messages: list[dict[str, str]], model: str | None = None) -> Iterator[str]: model = model or self.config.model_simple @@ -118,6 +136,9 @@ def _usage(response: Any) -> dict[str, int | None] | None: } def chat(self, messages: list[dict[str, str]], model: str | None = None) -> tuple[str, dict | None]: + return self.chat_response(messages, model).as_legacy_tuple() + + def chat_response(self, messages: list[dict[str, str]], model: str | None = None) -> ModelResponse: model = model or self.config.model_simple started = time.monotonic() @@ -127,7 +148,14 @@ def chat(self, messages: list[dict[str, str]], model: str | None = None) -> tupl usage = self._usage(response) usage_ledger._track_cost(self.config.provider, usage, model=getattr(response, "model", None) or model, latency_ms=round((time.monotonic() - started) * 1000)) - return str(getattr(response, "output_text", "") or "").strip(), usage + return model_response(provider=self.name, model=model, actual_model=getattr(response, "model", None), + text=str(getattr(response, "output_text", "") or ""), usage=usage, + calls=responses_output(response), response_id=getattr(response, "id", None), + raw_finish_reason=getattr(response, "status", None), + finish_detail=getattr(getattr(response, "incomplete_details", None), "reason", None), + blocked=any(getattr(part, "type", None) == "refusal" + for item in getattr(response, "output", None) or () + for part in getattr(item, "content", None) or ())) def chat_stream(self, messages: list[dict[str, str]], model: str | None = None) -> Iterator[str]: model = model or self.config.model_simple diff --git a/openkyrozen/providers/perplexity.py b/openkyrozen/providers/perplexity.py index 57a7cfb..ce3e174 100644 --- a/openkyrozen/providers/perplexity.py +++ b/openkyrozen/providers/perplexity.py @@ -6,6 +6,7 @@ import time from typing import Any, Iterator from openkyrozen.providers.base import LLMProvider +from openkyrozen.providers.models import ModelResponse, model_response, responses_output from openkyrozen.providers.config import ProviderConfig from openkyrozen.providers.retry import _retry_with_backoff @@ -43,6 +44,9 @@ def _usage(response: Any) -> dict[str, int | None] | None: } def chat(self, messages: list[dict[str, str]], model: str | None = None) -> tuple[str, dict | None]: + return self.chat_response(messages, model).as_legacy_tuple() + + def chat_response(self, messages: list[dict[str, str]], model: str | None = None) -> ModelResponse: model = model or self.config.model_simple prompt, instructions = self._prompt(messages) started = time.monotonic() @@ -53,7 +57,14 @@ def chat(self, messages: list[dict[str, str]], model: str | None = None) -> tupl usage = self._usage(response) usage_ledger._track_cost(self.config.provider, usage, model=model, latency_ms=round((time.monotonic() - started) * 1000)) - return str(getattr(response, "output_text", "") or "").strip(), usage + return model_response(provider=self.name, model=model, actual_model=getattr(response, "model", None), + text=str(getattr(response, "output_text", "") or ""), usage=usage, + calls=responses_output(response), response_id=getattr(response, "id", None), + raw_finish_reason=getattr(response, "status", None), + finish_detail=getattr(getattr(response, "incomplete_details", None), "reason", None), + blocked=any(getattr(part, "type", None) == "refusal" + for item in getattr(response, "output", None) or () + for part in getattr(item, "content", None) or ())) def chat_stream(self, messages: list[dict[str, str]], model: str | None = None) -> Iterator[str]: model = model or self.config.model_simple diff --git a/tests/test_provider_contracts.py b/tests/test_provider_contracts.py new file mode 100644 index 0000000..85635b1 --- /dev/null +++ b/tests/test_provider_contracts.py @@ -0,0 +1,388 @@ +"""Offline regression checks for provider-independent response boundaries.""" +import unittest +import copy +from types import SimpleNamespace as NS +from unittest.mock import Mock, patch + +from openkyrozen.providers import ProviderConfig +from openkyrozen.providers.openai import OpenAICompatProvider + +from openkyrozen.providers.models import FinishReason, ModelResponse, ToolCall, ProviderCapabilities, model_response +from openkyrozen.providers.base import LLMProvider, get_model_response +from openkyrozen.providers.openai import OpenAIResponsesProvider +from openkyrozen.providers.azure import AzureOpenAIProvider +from openkyrozen.providers.anthropic import AnthropicProvider +from openkyrozen.providers.google import GoogleProvider, VertexProvider +from openkyrozen.providers.bedrock import BedrockProvider +from openkyrozen.providers.ollama import OllamaNativeProvider +from openkyrozen.providers.perplexity import PerplexityProvider +from openkyrozen.providers.fallback import FallbackProvider + + +class ProviderResponseRegressionTests(unittest.TestCase): + def test_native_call_and_finish_are_preserved(self): + provider = OpenAICompatProvider.__new__(OpenAICompatProvider) + provider.config = ProviderConfig(provider="custom", model_simple="model") + provider._client = NS(chat=NS(completions=NS(create=Mock(return_value=NS( + id="response-id", model="model", usage=None, + choices=[NS(finish_reason="tool_calls", message=NS(content="", tool_calls=[ + NS(id="call-id", type="function", function=NS(name="read_file", arguments='{"path":"a"}')), + ]))], + ))))) + with patch("openkyrozen.providers.usage._track_cost"): + response = provider.chat_response([]) + self.assertEqual(response.tool_calls[0].id, "call-id") + self.assertEqual(response.tool_calls[0].arguments, {"path": "a"}) + self.assertEqual(response.finish_reason.value, "tool_request") + self.assertEqual(response.metadata["response_id"], "response-id") + + def test_legacy_conversion_refuses_length_stop(self): + provider = OpenAICompatProvider.__new__(OpenAICompatProvider) + provider.config = ProviderConfig(provider="custom", model_simple="model") + provider._client = NS(chat=NS(completions=NS(create=Mock(return_value=NS( + choices=[NS(finish_reason="length", message=NS(content="partial", tool_calls=[]))], usage=None, + ))))) + with patch("openkyrozen.providers.usage._track_cost"), self.assertRaises(ValueError): + provider.chat([]) + + def test_runtime_receives_typed_calls_and_rejects_text_conversion(self): + import tempfile + from pathlib import Path + from openkyrozen.app.bootstrap import build_application, build_memory + from openkyrozen.providers.models import FinishReason, ModelResponse, ToolCall + response = ModelResponse(tool_calls=(ToolCall("id", "read_file", {"path": "a"}),), + usage={"prompt_tokens": 2, "completion_tokens": 3}, + finish_reason=FinishReason.TOOL_REQUEST) + with tempfile.TemporaryDirectory() as directory, patch.dict("os.environ", {"KYROZEN_DISABLE_VECTOR_INDEX": "1"}): + application = build_application(memory=build_memory(Path(directory) / "state.sqlite3")) + runtime = application.runtime + runtime.llm_provider = NS(chat_response=Mock(return_value=response)) + try: + self.assertIs(runtime._get_model_response([]), response) + self.assertEqual(runtime._last_prompt_tokens, 2) + self.assertIn("text conversion refused", runtime._get_llm_response([])) + self.assertEqual(runtime.memory_bank.store.list_events("execution.receipt"), []) + finally: + application.close() + + + +class ContractTests(unittest.TestCase): + def test_call_validation_and_no_aliasing_of_sdk_arguments(self): + arguments = {"nested": [None, True, 12, {"unicode": "雪"}]} + call = ToolCall("id", "name", arguments) + arguments["nested"].append("changed") + self.assertEqual(call.arguments, {"nested": [None, True, 12, {"unicode": "雪"}]}) + cycle = {} + cycle["cycle"] = cycle + for value in (cycle, [], None, "{}", {1: "bad"}, {"x": float("nan")}, {"x": object()}, {"x": (1,)}): + with self.subTest(value=value), self.assertRaises(ValueError): + ToolCall("id", "name", value) + for call_id, name in (("", "name"), (None, "name"), ("id", " "), ("id", 7)): + with self.subTest(call_id=call_id, name=name), self.assertRaises(ValueError): + ToolCall(call_id, name, {}) + + def test_finish_mapping_and_checked_conversion(self): + cases = [("stop", FinishReason.FINAL), ("completed", FinishReason.FINAL), + ("end_turn", FinishReason.FINAL), ("stop_sequence", FinishReason.FINAL), + ("tool_calls", FinishReason.TOOL_REQUEST), ("tool_use", FinishReason.TOOL_REQUEST), + ("length", FinishReason.LENGTH), ("MAX_TOKENS", FinishReason.LENGTH), + ("failed", FinishReason.ERROR), ("cancelled", FinishReason.CANCELLED), + ("content_filter", FinishReason.BLOCKED), ("SAFETY", FinishReason.BLOCKED), + (None, FinishReason.UNKNOWN), ("new_reason", FinishReason.UNKNOWN)] + for raw, expected in cases: + with self.subTest(raw=raw): + response = model_response(provider="test", model="m", text=" text ", usage=None, raw_finish_reason=raw, + calls=[("id", "read", {})] if expected == FinishReason.TOOL_REQUEST else []) + self.assertEqual(response.finish_reason, expected) + self.assertEqual(response.metadata["raw_finish_reason"], raw) + if expected in {FinishReason.FINAL, FinishReason.UNKNOWN}: + self.assertEqual(response.as_legacy_tuple(), ("text", None)) + else: + with self.assertRaises(ValueError): + response.as_legacy_tuple() + for raw in ("incomplete", "in_progress", "queued", "pause_turn"): + with self.subTest(raw=raw), self.assertRaises(ValueError): + model_response(provider="test", model="m", text="partial", usage=None, + raw_finish_reason=raw).as_legacy_tuple() + response = model_response(provider="test", model="m", text="", usage=None, + raw_finish_reason="incomplete", finish_detail="max_output_tokens") + self.assertEqual(response.finish_reason, FinishReason.LENGTH) + + def test_canonical_ids_json_validation_and_usage_retention(self): + usage = {"prompt_tokens": 1, "prompt_cache_hit_tokens": 1, "reasoning_tokens": 2, "_estimated": 1} + response = model_response(provider="test", model="m", text=" raw ", usage=usage, + raw_finish_reason="stop", response_id="r", calls=[ + ("native", "one", '{"path":"a"}'), (None, "two", {}), + ]) + self.assertEqual(response.text, " raw ") + self.assertEqual(response.usage, usage) + self.assertEqual([call.id for call in response.tool_calls], ["native", "call_r_1"]) + self.assertEqual(response.metadata["synthesized_call_ids"], ["call_r_1"]) + self.assertEqual(response.finish_reason, FinishReason.TOOL_REQUEST) + for arguments in ('{"broken":', '[]', 'null', '{"a":1,"a":2}', '{"a":Infinity}'): + with self.subTest(arguments=arguments), self.assertRaises(ValueError): + model_response(provider="test", model="m", text="", usage=None, calls=[("id", "name", arguments)]) + with self.assertRaises(ValueError): + ModelResponse(tool_calls=None) + with self.assertRaises(ValueError): + ModelResponse(finish_reason=FinishReason.TOOL_REQUEST) + with self.assertRaises(ValueError): + ModelResponse(tool_calls=(ToolCall("id", "one", {}), ToolCall("id", "two", {}))) + other = model_response(provider="test", model="m", text="", usage=None, calls=[(None, "name", {})]) + another = model_response(provider="test", model="m", text="", usage=None, calls=[(None, "name", {})]) + self.assertNotEqual(other.tool_calls[0].id, another.tool_calls[0].id) + + def test_legacy_provider_bridge_and_capabilities(self): + class Legacy(LLMProvider): + def chat(self, messages, model=None): + return " legacy ", {"prompt_tokens": 1} + provider = Legacy(ProviderConfig(provider="custom", model_simple="m")) + response = provider.chat_response([]) + self.assertEqual(response.text, " legacy ") + self.assertEqual(response.finish_reason, FinishReason.UNKNOWN) + self.assertTrue(response.metadata["legacy"]) + self.assertEqual(response.as_legacy_tuple(), (" legacy ", {"prompt_tokens": 1})) + self.assertEqual(provider.get_capabilities(), ProviderCapabilities()) + self.assertEqual(get_model_response(NS(chat=lambda *args: ("duck", None)), []).text, "duck") + dynamic = Mock() + dynamic.chat.return_value = ("mock", None) + self.assertEqual(get_model_response(dynamic, []).text, "mock") + with self.assertRaises(ValueError): + get_model_response(NS(chat_response=lambda *args: ("invalid", None)), []) + with self.assertRaises(ValueError): + ProviderCapabilities(native_tools=1) + self.assertEqual(ProviderCapabilities.intersection([]), ProviderCapabilities()) + + +class AdapterContractTests(unittest.TestCase): + def adapter(self, cls, wire): + provider = cls.__new__(cls) + provider.config = ProviderConfig(provider="custom", model_simple="quick", model_complex="slow") + send = Mock(return_value=wire) + provider._client = NS(chat=NS(completions=NS(create=send)), responses=NS(create=send), + messages=NS(create=send), models=NS(generate_content=send), converse=send) + provider._base = "http://localhost:11434" + provider._requests = NS(post=Mock(return_value=NS(raise_for_status=lambda: None, json=lambda: wire))) + return provider + + def wires(self, text=" text ", calls=True): + tools = [NS(id="id1", function=NS(name="first", arguments='{"n":1}')), + NS(id="id2", function=NS(name="second", arguments="{}"))] if calls else [] + chat = NS(id="r", model="actual", choices=[NS(finish_reason="tool_calls" if calls else "stop", + message=NS(content=text, tool_calls=tools))], usage=None) + output = [NS(type="function_call", call_id=tool.id, id="item-" + tool.id, + name=tool.function.name, arguments=tool.function.arguments) for tool in tools] + responses = NS(id="r", output_text=text, output=output, status="completed", usage=None) + anthropic = NS(id="r", content=[NS(type="text", text=text)] + [ + NS(type="tool_use", id=tool.id, name=tool.function.name, input={"n": 1} if i == 0 else {}) + for i, tool in enumerate(tools)], stop_reason="tool_use" if calls else "end_turn", usage=None) + google = NS(response_id="r", candidates=[NS(finish_reason="STOP", content=NS(parts=[NS(text=text)] + [ + NS(function_call=NS(id=tool.id, name=tool.function.name, args={"n": 1} if i == 0 else {})) + for i, tool in enumerate(tools)]))], usage_metadata=None) + bedrock = {"output": {"message": {"content": [{"text": text}] + [ + {"toolUse": {"toolUseId": tool.id, "name": tool.function.name, "input": {"n": 1} if i == 0 else {}}} + for i, tool in enumerate(tools)]}}, "stopReason": "tool_use" if calls else "end_turn"} + ollama = {"message": {"content": text, "tool_calls": [ + {"id": tool.id, "function": {"name": tool.function.name, "arguments": {"n": 1} if i == 0 else {}}} + for i, tool in enumerate(tools)]}, "done_reason": "stop"} + return [(OpenAICompatProvider, chat), (AzureOpenAIProvider, chat), + (OpenAIResponsesProvider, responses), (PerplexityProvider, responses), + (AnthropicProvider, anthropic), (GoogleProvider, google), (VertexProvider, google), + (BedrockProvider, bedrock), (OllamaNativeProvider, ollama)] + + def test_all_adapter_families_preserve_calls_text_and_charge_once(self): + for text, calls in ((" text ", True), ("", True), (" text ", False)): + for cls, wire in self.wires(text, calls): + with self.subTest(cls=cls.__name__, text=text, calls=calls), patch("openkyrozen.providers.usage._track_cost") as charge: + provider = self.adapter(cls, wire) + response = provider.chat_response([]) + self.assertIsInstance(response, ModelResponse) + self.assertEqual(response.text, text) + self.assertEqual([call.id for call in response.tool_calls], ["id1", "id2"] if calls else []) + self.assertEqual(response.finish_reason, FinishReason.TOOL_REQUEST if calls else FinishReason.FINAL) + self.assertEqual(response.metadata["provider"], "custom") + charge.assert_called_once() + charge.reset_mock() + if calls: + with self.assertRaises(ValueError): + provider.chat([]) + else: + self.assertEqual(provider.chat([])[0], text.strip()) + charge.assert_called_once() + caps = provider.get_capabilities("quick") + self.assertFalse(caps.native_tools) + self.assertFalse(caps.streaming_tool_calls) + self.assertEqual(caps.text_streaming, cls is not OllamaNativeProvider) + + def test_missing_ids_and_malformed_args_across_transports(self): + for cls, wire in self.wires(): + wire = copy.deepcopy(wire) + if isinstance(wire, dict): + tool = (wire["message"]["tool_calls"][0] if cls is OllamaNativeProvider + else wire["output"]["message"]["content"][1]["toolUse"]) + id_key = "id" if cls is OllamaNativeProvider else "toolUseId" + tool.pop(id_key) + target = tool["function"] if cls is OllamaNativeProvider else tool + arg_key = "arguments" if cls is OllamaNativeProvider else "input" + elif cls in {OpenAICompatProvider, AzureOpenAIProvider}: + tool = wire.choices[0].message.tool_calls[0]; tool.id = None + target, arg_key = tool.function, "arguments" + elif cls in {OpenAIResponsesProvider, PerplexityProvider}: + tool = wire.output[0]; tool.call_id = None + target, arg_key = tool, "arguments" + elif cls is AnthropicProvider: + tool = wire.content[1]; tool.id = None + target, arg_key = tool, "input" + else: + tool = wire.candidates[0].content.parts[1].function_call; tool.id = None + target, arg_key = tool, "args" + with self.subTest(cls=cls.__name__), patch("openkyrozen.providers.usage._track_cost"): + response = self.adapter(cls, wire).chat_response([]) + self.assertIn(response.tool_calls[0].id, response.metadata["synthesized_call_ids"]) + if isinstance(target, dict): + target[arg_key] = [] + else: + setattr(target, arg_key, []) + with self.assertRaises(ValueError): + self.adapter(cls, wire).chat_response([]) + + def test_sdk_openai_payloads_use_call_id_not_output_item_id(self): + from openai.types.chat import ChatCompletion + from openai.types.responses import Response + chat = ChatCompletion.model_validate({"id": "r", "object": "chat.completion", "created": 1, + "model": "actual", "choices": [{"index": 0, "finish_reason": "tool_calls", "message": { + "role": "assistant", "content": None, "tool_calls": [{"id": "call", "type": "function", + "function": {"name": "read", "arguments": "{}"}}]}}]}) + response = Response.model_construct(id="r", model="actual", status="completed", usage=None, + output=[{"type": "function_call", "call_id": "call", "id": "item", "name": "read", "arguments": "{}"}]) + # Validate the output item with the SDK rather than relying on dict attribute access. + from openai.types.responses import ResponseFunctionToolCall + response.output = [ResponseFunctionToolCall.model_validate(response.output[0])] + with patch("openkyrozen.providers.usage._track_cost"): + parsed = self.adapter(OpenAICompatProvider, chat).chat_response([]) + self.assertEqual(parsed.tool_calls[0].id, "call") + self.assertEqual(parsed.metadata["model"], "actual") + self.assertEqual(self.adapter(OpenAIResponsesProvider, response).chat_response([]).tool_calls[0].id, "call") + + def test_fallback_preserves_identity_and_does_not_replay_conversion_failure(self): + class Probe(LLMProvider): + def __init__(self, config, response, caps): + super().__init__(config); self.response = response; self.caps = caps; self.models = [] + def chat(self, messages, model=None): + return self.chat_response(messages, model).as_legacy_tuple() + def chat_response(self, messages, model=None): + self.models.append(model) + if isinstance(self.response, Exception): + raise self.response + return self.response + def get_capabilities(self, model=None): + return self.caps + result = ModelResponse(tool_calls=(ToolCall("id", "read", {}),), finish_reason=FinishReason.TOOL_REQUEST, + metadata={"provider": "fallback", "model": "f-slow"}) + primary = Probe(ProviderConfig(provider="custom", model_simple="p-quick", model_complex="p-slow"), + RuntimeError("offline failure"), ProviderCapabilities(text_streaming=True, native_tools=True)) + fallback = Probe(ProviderConfig(provider="custom", model_simple="f-quick", model_complex="f-slow"), + result, ProviderCapabilities(text_streaming=True)) + wrapper = FallbackProvider.__new__(FallbackProvider) + wrapper._primary, wrapper._fallbacks = primary, [fallback] + self.assertIs(wrapper.chat_response([], "p-slow"), result) + self.assertEqual(fallback.models, ["f-slow"]) + self.assertEqual(wrapper.get_capabilities(), ProviderCapabilities(text_streaming=True)) + primary.response = result + with self.assertRaises(ValueError): + wrapper.chat([], "p-slow") + self.assertEqual(fallback.models, ["f-slow"]) + + def test_perplexity_sdk_message_output_is_not_lost(self): + try: + from perplexity.types import ResponseCreateResponse + except ImportError: + self.skipTest("optional Perplexity SDK unavailable") + payload = ResponseCreateResponse.model_validate({ + "background": False, "created_at": 1, "error": None, "id": "r", "model": "actual", + "object": "response", "previous_response_id": None, "status": "completed", "store": False, + "usage": None, "output": [ + {"id": "msg", "type": "message", "role": "assistant", "status": "completed", + "content": [{"type": "output_text", "text": " sdk text ", "annotations": []}]}, + {"id": "item", "type": "function_call", "call_id": "native", "name": "read", + "arguments": "{}", "status": "completed"}, + ], + }) + with patch("openkyrozen.providers.usage._track_cost"): + response = self.adapter(PerplexityProvider, payload).chat_response([]) + self.assertEqual(response.text, " sdk text ") + self.assertEqual(response.tool_calls[0].id, "native") + + def test_other_installed_sdk_response_shapes(self): + try: + from anthropic.types import Message + from google.genai.types import GenerateContentResponse + from botocore.loaders import Loader + from botocore.model import ServiceModel + from botocore.validate import validate_parameters + except ImportError: + self.skipTest("optional native SDKs unavailable") + anthropic = Message.model_validate({"id": "msg", "type": "message", "role": "assistant", "model": "claude", + "content": [{"type": "text", "text": "text"}, {"type": "tool_use", "id": "native", "name": "read", "input": {}}], + "stop_reason": "tool_use", "stop_sequence": None, "usage": {"input_tokens": 2, "output_tokens": 3}}) + google = GenerateContentResponse.model_validate({"candidates": [{"finish_reason": "STOP", "content": { + "role": "model", "parts": [{"text": "text"}, {"function_call": {"id": "native", "name": "read", "args": {}}}] + }}], "model_version": "gemini", "usage_metadata": {"prompt_token_count": 2, "candidates_token_count": 3}}) + bedrock = {"output": {"message": {"role": "assistant", "content": [{"text": "text"}, + {"toolUse": {"toolUseId": "native", "name": "read", "input": {}}}]}}, "stopReason": "tool_use", + "usage": {"inputTokens": 2, "outputTokens": 3, "totalTokens": 5}, "metrics": {"latencyMs": 1}} + model = ServiceModel(Loader().load_service_model("bedrock-runtime", "service-2")) + validate_parameters(bedrock, model.operation_model("Converse").output_shape) + with patch("openkyrozen.providers.usage._track_cost"): + for cls, payload in ((AnthropicProvider, anthropic), (GoogleProvider, google), (BedrockProvider, bedrock)): + with self.subTest(cls=cls.__name__): + response = self.adapter(cls, payload).chat_response([]) + self.assertEqual(response.tool_calls[0].id, "native") + self.assertEqual(response.text, "text") + self.assertEqual(response.usage["prompt_tokens"], 2) + + def test_transport_finish_statuses_and_unknown_statuses(self): + for cls, wire in self.wires(calls=False): + cases = [(None, FinishReason.UNKNOWN)] + if cls in {OpenAICompatProvider, AzureOpenAIProvider}: + cases += [("length", FinishReason.LENGTH), ("content_filter", FinishReason.BLOCKED)] + elif cls in {OpenAIResponsesProvider, PerplexityProvider}: + cases += [("failed", FinishReason.ERROR), ("cancelled", FinishReason.CANCELLED), ("incomplete", FinishReason.UNKNOWN)] + elif cls is AnthropicProvider: + cases += [("max_tokens", FinishReason.LENGTH), ("refusal", FinishReason.BLOCKED)] + elif cls in {GoogleProvider, VertexProvider}: + cases += [("MAX_TOKENS", FinishReason.LENGTH), ("SAFETY", FinishReason.BLOCKED), ("MALFORMED_FUNCTION_CALL", FinishReason.ERROR)] + elif cls is BedrockProvider: + cases += [("max_tokens", FinishReason.LENGTH), ("guardrail_intervened", FinishReason.BLOCKED)] + else: + cases += [("length", FinishReason.LENGTH), ("error", FinishReason.ERROR)] + for raw, expected in cases: + payload = copy.deepcopy(wire) + if cls in {OpenAICompatProvider, AzureOpenAIProvider}: + payload.choices[0].finish_reason = raw + elif cls in {OpenAIResponsesProvider, PerplexityProvider}: + payload.status = raw + elif cls is AnthropicProvider: + payload.stop_reason = raw + elif cls in {GoogleProvider, VertexProvider}: + payload.candidates[0].finish_reason = raw + else: + payload["stopReason" if cls is BedrockProvider else "done_reason"] = raw + with self.subTest(cls=cls.__name__, raw=raw), patch("openkyrozen.providers.usage._track_cost"): + response = self.adapter(cls, payload).chat_response([]) + self.assertEqual(response.finish_reason, expected) + self.assertEqual(response.metadata["raw_finish_reason"], raw) + + def test_malformed_native_response_is_not_replayed_on_fallback(self): + primary = self.adapter(OpenAICompatProvider, self.wires()[0][1]) + primary._client.chat.completions.create.return_value.choices[0].message.tool_calls[0].function.arguments = "[]" + fallback = NS(config=ProviderConfig(provider="custom", model_simple="f", model_complex="f"), + name="fallback", chat_response=Mock(return_value=ModelResponse(text="fallback"))) + wrapper = FallbackProvider.__new__(FallbackProvider) + wrapper._primary, wrapper._fallbacks = primary, [fallback] + with patch("openkyrozen.providers.usage._track_cost") as charge, self.assertRaises(ValueError): + wrapper.chat_response([]) + primary._client.chat.completions.create.assert_called_once() + fallback.chat_response.assert_not_called() + charge.assert_called_once() From b720ebf5b0e4641c06ca7c0cb75717d400bf1c36 Mon Sep 17 00:00:00 2001 From: Evan Date: Thu, 8 Oct 2026 19:59:02 +0800 Subject: [PATCH 2/2] fix: fail closed on received provider response errors --- docs/providers.md | 4 +- openkyrozen/providers/anthropic.py | 39 ++++++++-------- openkyrozen/providers/bedrock.py | 23 +++++----- openkyrozen/providers/google.py | 45 ++++++++++--------- openkyrozen/providers/models.py | 22 +++++++-- openkyrozen/providers/ollama.py | 41 ++++++++++------- openkyrozen/providers/openai.py | 70 +++++++++++++++-------------- openkyrozen/providers/perplexity.py | 25 ++++++----- tests/test_provider_contracts.py | 52 ++++++++++++++++++++- 9 files changed, 203 insertions(+), 118 deletions(-) diff --git a/docs/providers.md b/docs/providers.md index 6001e09..6fa9a99 100644 --- a/docs/providers.md +++ b/docs/providers.md @@ -58,7 +58,9 @@ The registry marks a subset of providers for automatic selection and records pre `text`, ordered `ToolCall` objects, normalized `finish_reason`, usage and explicit provider/model/response metadata. `ToolCall` preserves a native call ID, name and JSON-object arguments. Contract validation raises `ProviderContractError`, a -`ValueError` subclass; fallback does not replay a received malformed response. IDs omitted by a transport are generated within the response +`ValueError` subclass; fallback does not replay a received malformed response. +All processing after a successful SDK return is marked as the received-response +phase, so parsing or accounting failures cannot trigger another generation. IDs omitted by a transport are generated within the response scope and listed in `metadata["synthesized_call_ids"]`. Malformed arguments, duplicate IDs/JSON keys and non-JSON values are rejected; calls are never rendered as assistant text or executed by the provider boundary. diff --git a/openkyrozen/providers/anthropic.py b/openkyrozen/providers/anthropic.py index 21ffb65..64134b0 100644 --- a/openkyrozen/providers/anthropic.py +++ b/openkyrozen/providers/anthropic.py @@ -5,7 +5,7 @@ import time from typing import Any, Iterator from openkyrozen.providers.base import LLMProvider -from openkyrozen.providers.models import ModelResponse, model_response +from openkyrozen.providers.models import received_response, ModelResponse, model_response from openkyrozen.providers.config import ProviderConfig from openkyrozen.providers.retry import _retry_with_backoff @@ -59,24 +59,25 @@ def _call(): return self._client.messages.create(**kwargs) response = _retry_with_backoff(_call) - text = "" - for block in response.content: - if hasattr(block, "text"): - text += block.text - usage = getattr(response, "usage", None) - usage_dict = None - if usage is not None: - usage_dict = { - "prompt_tokens": getattr(usage, "input_tokens", 0) or 0, - "completion_tokens": getattr(usage, "output_tokens", 0) or 0, - } - usage_ledger._track_cost(self.config.provider, usage_dict, model=getattr(response, "model", None) or model, - latency_ms=round((time.monotonic() - started) * 1000)) - calls = [(getattr(block, "id", None), getattr(block, "name", None), getattr(block, "input", None)) - for block in response.content if getattr(block, "type", None) == "tool_use"] - return model_response(provider=self.name, model=model, actual_model=getattr(response, "model", None), text=text, usage=usage_dict, - calls=calls, response_id=getattr(response, "id", None), - raw_finish_reason=getattr(response, "stop_reason", None)) + with received_response(): + usage = getattr(response, "usage", None) + usage_dict = None + if usage is not None: + usage_dict = { + "prompt_tokens": getattr(usage, "input_tokens", 0) or 0, + "completion_tokens": getattr(usage, "output_tokens", 0) or 0, + } + usage_ledger._track_cost(self.config.provider, usage_dict, model=getattr(response, "model", None) or model, + latency_ms=round((time.monotonic() - started) * 1000)) + text = "" + for block in response.content: + if hasattr(block, "text"): + text += block.text + calls = [(getattr(block, "id", None), getattr(block, "name", None), getattr(block, "input", None)) + for block in response.content if getattr(block, "type", None) == "tool_use"] + return model_response(provider=self.name, model=model, actual_model=getattr(response, "model", None), text=text, usage=usage_dict, + calls=calls, response_id=getattr(response, "id", None), + raw_finish_reason=getattr(response, "stop_reason", None)) def chat_stream(self, messages: list[dict[str, str]], model: str | None = None) -> Iterator[str]: model = model or self.config.model_simple diff --git a/openkyrozen/providers/bedrock.py b/openkyrozen/providers/bedrock.py index eb20115..9f76a90 100644 --- a/openkyrozen/providers/bedrock.py +++ b/openkyrozen/providers/bedrock.py @@ -6,7 +6,7 @@ import time from typing import Any, Iterator from openkyrozen.providers.base import LLMProvider -from openkyrozen.providers.models import ModelResponse, model_response +from openkyrozen.providers.models import received_response, ModelResponse, model_response from openkyrozen.providers.config import ProviderConfig from openkyrozen.providers.retry import _retry_with_backoff @@ -64,16 +64,17 @@ def chat_response(self, messages: list[dict[str, str]], model: str | None = None if system: kwargs["system"] = system response = _retry_with_backoff(lambda: self._client.converse(**kwargs)) - content = response.get("output", {}).get("message", {}).get("content", []) - text = "".join(str(item.get("text", "")) for item in content if isinstance(item, dict)) - usage = self._usage(response) - usage_ledger._track_cost(self.config.provider, usage, model=model, - latency_ms=round((time.monotonic() - started) * 1000)) - calls = [(item["toolUse"].get("toolUseId"), item["toolUse"].get("name"), item["toolUse"].get("input")) - for item in content if isinstance(item, dict) and "toolUse" in item] - return model_response(provider=self.name, model=model, text=text, usage=usage, - calls=calls, response_id=response.get("ResponseMetadata", {}).get("RequestId"), - raw_finish_reason=response.get("stopReason")) + with received_response(): + usage = self._usage(response) + usage_ledger._track_cost(self.config.provider, usage, model=model, + latency_ms=round((time.monotonic() - started) * 1000)) + content = response.get("output", {}).get("message", {}).get("content", []) + text = "".join(str(item.get("text", "")) for item in content if isinstance(item, dict)) + calls = [(item["toolUse"].get("toolUseId"), item["toolUse"].get("name"), item["toolUse"].get("input")) + for item in content if isinstance(item, dict) and "toolUse" in item] + return model_response(provider=self.name, model=model, text=text, usage=usage, + calls=calls, response_id=response.get("ResponseMetadata", {}).get("RequestId"), + raw_finish_reason=response.get("stopReason")) def chat_stream(self, messages: list[dict[str, str]], model: str | None = None) -> Iterator[str]: model = model or self.config.model_simple diff --git a/openkyrozen/providers/google.py b/openkyrozen/providers/google.py index e53db82..7f95949 100644 --- a/openkyrozen/providers/google.py +++ b/openkyrozen/providers/google.py @@ -6,7 +6,7 @@ import time from typing import Any, Iterator from openkyrozen.providers.base import LLMProvider -from openkyrozen.providers.models import ModelResponse, model_response +from openkyrozen.providers.models import received_response, ModelResponse, model_response from openkyrozen.providers.config import ProviderConfig from openkyrozen.providers.retry import _retry_with_backoff @@ -71,27 +71,28 @@ def _call(): ) response = _retry_with_backoff(_call) - candidates = getattr(response, "candidates", None) or () - candidate = candidates[0] if candidates else None - parts = getattr(getattr(candidate, "content", None), "parts", None) - text = ("".join(part.text for part in parts if getattr(part, "text", None) - and not getattr(part, "thought", False)) if parts is not None - else str(getattr(response, "text", "") or "")) - usage_dict = self._usage(response) - usage_ledger._track_cost(self.config.provider, usage_dict, model=model, - latency_ms=round((time.monotonic() - started) * 1000)) - calls = [] - for part in parts or (): - call = getattr(part, "function_call", None) - if call is not None: - arguments = getattr(call, "args", None) - calls.append((getattr(call, "id", None), getattr(call, "name", None), - {} if arguments is None else arguments)) - block_reason = getattr(getattr(response, "prompt_feedback", None), "block_reason", None) - return model_response(provider=self.name, model=model, actual_model=getattr(response, "model_version", None), text=text, usage=usage_dict, - calls=calls, response_id=getattr(response, "response_id", None), - raw_finish_reason=getattr(candidate, "finish_reason", None) or block_reason, - blocked=block_reason not in {None, "BLOCK_REASON_UNSPECIFIED"}) + with received_response(): + usage_dict = self._usage(response) + usage_ledger._track_cost(self.config.provider, usage_dict, model=model, + latency_ms=round((time.monotonic() - started) * 1000)) + candidates = getattr(response, "candidates", None) or () + candidate = candidates[0] if candidates else None + parts = getattr(getattr(candidate, "content", None), "parts", None) + text = ("".join(part.text for part in parts if getattr(part, "text", None) + and not getattr(part, "thought", False)) if parts is not None + else str(getattr(response, "text", "") or "")) + calls = [] + for part in parts or (): + call = getattr(part, "function_call", None) + if call is not None: + arguments = getattr(call, "args", None) + calls.append((getattr(call, "id", None), getattr(call, "name", None), + {} if arguments is None else arguments)) + block_reason = getattr(getattr(response, "prompt_feedback", None), "block_reason", None) + return model_response(provider=self.name, model=model, actual_model=getattr(response, "model_version", None), text=text, usage=usage_dict, + calls=calls, response_id=getattr(response, "response_id", None), + raw_finish_reason=getattr(candidate, "finish_reason", None) or block_reason, + blocked=block_reason not in {None, "BLOCK_REASON_UNSPECIFIED", "BLOCKED_REASON_UNSPECIFIED"}) def chat_stream(self, messages: list[dict[str, str]], model: str | None = None) -> Iterator[str]: model = model or self.config.model_simple diff --git a/openkyrozen/providers/models.py b/openkyrozen/providers/models.py index f7aa4bd..1e8fe45 100644 --- a/openkyrozen/providers/models.py +++ b/openkyrozen/providers/models.py @@ -6,6 +6,7 @@ import math import uuid from dataclasses import dataclass, field, fields +from contextlib import contextmanager from enum import StrEnum from typing import Any @@ -14,6 +15,17 @@ class ProviderContractError(ValueError): """A received response cannot satisfy the contract; do not replay generation.""" +@contextmanager +def received_response(): + """A successful SDK return ends failover eligibility, including parse/billing errors.""" + try: + yield + except ProviderContractError: + raise + except Exception as exc: + raise ProviderContractError("Received provider response could not be processed") from exc + + class FinishReason(StrEnum): FINAL = "final" TOOL_REQUEST = "tool_request" @@ -119,8 +131,10 @@ def __post_init__(self): def as_legacy_tuple(self) -> tuple[str, dict | None]: """Refuse to hide calls or unsuccessful completion in a text-only API.""" + raw = self.metadata.get("raw_finish_reason") + unfinished = isinstance(raw, str) and raw.lower() in ("incomplete", "in_progress", "queued", "pause_turn", "continuation") if (self.tool_calls or self.finish_reason not in {FinishReason.FINAL, FinishReason.UNKNOWN} - or self.metadata.get("raw_finish_reason") in ("incomplete", "in_progress", "queued", "pause_turn")): + or unfinished): raise ProviderContractError(f"Structured response ({self.finish_reason}) requires chat_response(); text conversion refused") return (self.text if self.metadata.get("legacy") else self.text.strip()), self.usage @@ -136,12 +150,14 @@ def normalize_finish(reason, *, has_calls=False, detail=None) -> FinishReason: return FinishReason.TOOL_REQUEST if reason in {"length", "max_tokens", "max_output_tokens", "max_token", "model_context_window_exceeded"}: return FinishReason.LENGTH - if reason in {"failed", "error", "malformed_function_call", "unexpected_tool_call"}: + if reason in {"failed", "error", "malformed_function_call", "unexpected_tool_call", + "too_many_tool_calls", "language", "no_image"}: return FinishReason.ERROR if reason in {"cancelled", "canceled"}: return FinishReason.CANCELLED if reason in {"content_filter", "content_filtered", "guardrail_intervened", "safety", "recitation", - "refusal", "blocked", "blocklist", "prohibited_content", "spii"}: + "refusal", "blocked", "blocklist", "prohibited_content", "spii", "model_armor", + "image_safety", "image_prohibited_content", "image_recitation"}: return FinishReason.BLOCKED return FinishReason.UNKNOWN diff --git a/openkyrozen/providers/ollama.py b/openkyrozen/providers/ollama.py index c5ef254..c6fba3d 100644 --- a/openkyrozen/providers/ollama.py +++ b/openkyrozen/providers/ollama.py @@ -4,7 +4,7 @@ import sys import time from openkyrozen.providers.base import LLMProvider -from openkyrozen.providers.models import ModelResponse, model_response +from openkyrozen.providers.models import received_response, ModelResponse, model_response from openkyrozen.providers.config import ProviderConfig @@ -29,17 +29,28 @@ def chat_response(self, messages: list[dict[str, str]], model: str | None = None payload = {"model": model, "messages": messages, "stream": False, "think": False} started = time.monotonic() resp = self._requests.post(url, json=payload, timeout=120) - resp.raise_for_status() - data = resp.json() - text = data.get("message", {}).get("content", "") - usage_dict = { - "prompt_tokens": data.get("prompt_eval_count", 0) or 0, - "completion_tokens": data.get("eval_count", 0) or 0, - } - usage_ledger._track_cost("ollama", usage_dict, model=model, - latency_ms=round((time.monotonic() - started) * 1000)) - calls = [(call.get("id"), call.get("function", {}).get("name"), - call.get("function", {}).get("arguments", {})) - for call in data.get("message", {}).get("tool_calls", [])] - return model_response(provider=self.name, model=model, text=text, usage=usage_dict, - calls=calls, raw_finish_reason="error" if data.get("error") else data.get("done_reason")) + try: + resp.raise_for_status() + except self._requests.HTTPError as exc: + try: + body = resp.json() + except ValueError: + body = None + detail = body.get("error") if isinstance(body, dict) else None + if isinstance(detail, str): + raise self._requests.HTTPError(f"Ollama HTTP {resp.status_code}: {detail[:2000]}", response=resp) from exc + raise + with received_response(): + data = resp.json() + usage_dict = { + "prompt_tokens": data.get("prompt_eval_count", 0) or 0, + "completion_tokens": data.get("eval_count", 0) or 0, + } + usage_ledger._track_cost("ollama", usage_dict, model=model, + latency_ms=round((time.monotonic() - started) * 1000)) + text = data.get("message", {}).get("content", "") + calls = [(call.get("id"), call.get("function", {}).get("name"), + call.get("function", {}).get("arguments", {})) + for call in data.get("message", {}).get("tool_calls", [])] + return model_response(provider=self.name, model=model, text=text, usage=usage_dict, + calls=calls, raw_finish_reason="error" if data.get("error") else data.get("done_reason")) diff --git a/openkyrozen/providers/openai.py b/openkyrozen/providers/openai.py index 62ba343..693e088 100644 --- a/openkyrozen/providers/openai.py +++ b/openkyrozen/providers/openai.py @@ -5,7 +5,7 @@ import time from typing import Any, Iterator from openkyrozen.providers.base import LLMProvider -from openkyrozen.providers.models import ModelResponse, model_response, responses_output, ProviderContractError +from openkyrozen.providers.models import received_response, ModelResponse, model_response, responses_output, ProviderContractError from openkyrozen.providers.config import ProviderConfig, _provider_env_key from openkyrozen.providers.usage import _openai_usage_dict from openkyrozen.providers.retry import _retry_with_backoff @@ -41,28 +41,29 @@ def _call(): return response response = _retry_with_backoff(_call) - text = response.choices[0].message.content or "" - usage = getattr(response, "usage", None) - usage_dict = None - if usage is not None: - usage_dict = _openai_usage_dict(usage) - usage_ledger._track_cost(self.config.provider, usage_dict, model=getattr(response, "model", None) or model, - latency_ms=round((time.monotonic() - started) * 1000)) - message = response.choices[0].message - calls = [] - for call in getattr(message, "tool_calls", None) or (): - function = getattr(call, "function", None) - if function is None: - raise ProviderContractError("Unsupported native tool call type") - calls.append((getattr(call, "id", None), getattr(function, "name", None), - getattr(function, "arguments", None))) - legacy_call = getattr(message, "function_call", None) - if legacy_call is not None and not calls: - calls.append((None, getattr(legacy_call, "name", None), getattr(legacy_call, "arguments", None))) - return model_response(provider=self.name, model=model, actual_model=getattr(response, "model", None), text=text, usage=usage_dict, - calls=calls, response_id=getattr(response, "id", None), - raw_finish_reason=getattr(response.choices[0], "finish_reason", None), - blocked=bool(getattr(message, "refusal", None))) + with received_response(): + usage = getattr(response, "usage", None) + usage_dict = None + if usage is not None: + usage_dict = _openai_usage_dict(usage) + usage_ledger._track_cost(self.config.provider, usage_dict, model=getattr(response, "model", None) or model, + latency_ms=round((time.monotonic() - started) * 1000)) + message = response.choices[0].message + text = message.content or "" + calls = [] + for call in getattr(message, "tool_calls", None) or (): + function = getattr(call, "function", None) + if function is None: + raise ProviderContractError("Unsupported native tool call type") + calls.append((getattr(call, "id", None), getattr(function, "name", None), + getattr(function, "arguments", None))) + legacy_call = getattr(message, "function_call", None) + if legacy_call is not None and not calls: + calls.append((None, getattr(legacy_call, "name", None), getattr(legacy_call, "arguments", None))) + return model_response(provider=self.name, model=model, actual_model=getattr(response, "model", None), text=text, usage=usage_dict, + calls=calls, response_id=getattr(response, "id", None), + raw_finish_reason=getattr(response.choices[0], "finish_reason", None), + blocked=bool(getattr(message, "refusal", None))) def chat_stream(self, messages: list[dict[str, str]], model: str | None = None) -> Iterator[str]: model = model or self.config.model_simple @@ -145,17 +146,18 @@ def chat_response(self, messages: list[dict[str, str]], model: str | None = None response = _retry_with_backoff( lambda: self._client.responses.create(model=model, input=messages), ) - usage = self._usage(response) - usage_ledger._track_cost(self.config.provider, usage, model=getattr(response, "model", None) or model, - latency_ms=round((time.monotonic() - started) * 1000)) - return model_response(provider=self.name, model=model, actual_model=getattr(response, "model", None), - text=str(getattr(response, "output_text", "") or ""), usage=usage, - calls=responses_output(response), response_id=getattr(response, "id", None), - raw_finish_reason=getattr(response, "status", None), - finish_detail=getattr(getattr(response, "incomplete_details", None), "reason", None), - blocked=any(getattr(part, "type", None) == "refusal" - for item in getattr(response, "output", None) or () - for part in getattr(item, "content", None) or ())) + with received_response(): + usage = self._usage(response) + usage_ledger._track_cost(self.config.provider, usage, model=getattr(response, "model", None) or model, + latency_ms=round((time.monotonic() - started) * 1000)) + return model_response(provider=self.name, model=model, actual_model=getattr(response, "model", None), + text=str(getattr(response, "output_text", "") or ""), usage=usage, + calls=responses_output(response), response_id=getattr(response, "id", None), + raw_finish_reason=getattr(response, "status", None), + finish_detail=getattr(getattr(response, "incomplete_details", None), "reason", None), + blocked=any(getattr(part, "type", None) == "refusal" + for item in getattr(response, "output", None) or () + for part in getattr(item, "content", None) or ())) def chat_stream(self, messages: list[dict[str, str]], model: str | None = None) -> Iterator[str]: model = model or self.config.model_simple diff --git a/openkyrozen/providers/perplexity.py b/openkyrozen/providers/perplexity.py index ce3e174..7049274 100644 --- a/openkyrozen/providers/perplexity.py +++ b/openkyrozen/providers/perplexity.py @@ -6,7 +6,7 @@ import time from typing import Any, Iterator from openkyrozen.providers.base import LLMProvider -from openkyrozen.providers.models import ModelResponse, model_response, responses_output +from openkyrozen.providers.models import received_response, ModelResponse, model_response, responses_output from openkyrozen.providers.config import ProviderConfig from openkyrozen.providers.retry import _retry_with_backoff @@ -54,17 +54,18 @@ def chat_response(self, messages: list[dict[str, str]], model: str | None = None if instructions: kwargs["instructions"] = instructions response = _retry_with_backoff(lambda: self._client.responses.create(**kwargs)) - usage = self._usage(response) - usage_ledger._track_cost(self.config.provider, usage, model=model, - latency_ms=round((time.monotonic() - started) * 1000)) - return model_response(provider=self.name, model=model, actual_model=getattr(response, "model", None), - text=str(getattr(response, "output_text", "") or ""), usage=usage, - calls=responses_output(response), response_id=getattr(response, "id", None), - raw_finish_reason=getattr(response, "status", None), - finish_detail=getattr(getattr(response, "incomplete_details", None), "reason", None), - blocked=any(getattr(part, "type", None) == "refusal" - for item in getattr(response, "output", None) or () - for part in getattr(item, "content", None) or ())) + with received_response(): + usage = self._usage(response) + usage_ledger._track_cost(self.config.provider, usage, model=model, + latency_ms=round((time.monotonic() - started) * 1000)) + return model_response(provider=self.name, model=model, actual_model=getattr(response, "model", None), + text=str(getattr(response, "output_text", "") or ""), usage=usage, + calls=responses_output(response), response_id=getattr(response, "id", None), + raw_finish_reason=getattr(response, "status", None), + finish_detail=getattr(getattr(response, "incomplete_details", None), "reason", None), + blocked=any(getattr(part, "type", None) == "refusal" + for item in getattr(response, "output", None) or () + for part in getattr(item, "content", None) or ())) def chat_stream(self, messages: list[dict[str, str]], model: str | None = None) -> Iterator[str]: model = model or self.config.model_simple diff --git a/tests/test_provider_contracts.py b/tests/test_provider_contracts.py index 85635b1..fb80bfd 100644 --- a/tests/test_provider_contracts.py +++ b/tests/test_provider_contracts.py @@ -101,7 +101,7 @@ def test_finish_mapping_and_checked_conversion(self): else: with self.assertRaises(ValueError): response.as_legacy_tuple() - for raw in ("incomplete", "in_progress", "queued", "pause_turn"): + for raw in ("incomplete", "in_progress", "queued", "pause_turn", "CONTINUATION"): with self.subTest(raw=raw), self.assertRaises(ValueError): model_response(provider="test", model="m", text="partial", usage=None, raw_finish_reason=raw).as_legacy_tuple() @@ -386,3 +386,53 @@ def test_malformed_native_response_is_not_replayed_on_fallback(self): primary._client.chat.completions.create.assert_called_once() fallback.chat_response.assert_not_called() charge.assert_called_once() + + def test_all_received_response_parsing_failures_stop_fallback(self): + cases = self.wires(calls=False) + for cls, wire in cases: + wire = copy.deepcopy(wire) + if cls in {OpenAICompatProvider, AzureOpenAIProvider}: + wire.choices = [] + elif cls in {OpenAIResponsesProvider, PerplexityProvider}: + wire.output = [NS(type="function_call", call_id="id", name="read", arguments=[]) ] + elif cls is AnthropicProvider: + wire.content = [NS(type="text", text=7)] + elif cls in {GoogleProvider, VertexProvider}: + wire.candidates[0].content.parts = [NS(text=7)] + elif cls is BedrockProvider: + wire["output"]["message"]["content"] = [{"toolUse": []}] + else: + wire["message"] = [] + primary = self.adapter(cls, wire) + fallback = NS(config=ProviderConfig(provider="custom", model_simple="f", model_complex="f"), + name="fallback", chat_response=Mock(return_value=ModelResponse(text="fallback"))) + wrapper = FallbackProvider.__new__(FallbackProvider) + wrapper._primary, wrapper._fallbacks = primary, [fallback] + with self.subTest(cls=cls.__name__), patch("openkyrozen.providers.usage._track_cost"), self.assertRaises(ValueError): + wrapper.chat_response([]) + fallback.chat_response.assert_not_called() + + def test_google_sdk_blocked_and_unsuccessful_finish_reasons(self): + from openkyrozen.providers.models import normalize_finish + for raw in ("MODEL_ARMOR", "IMAGE_SAFETY", "IMAGE_PROHIBITED_CONTENT", "IMAGE_RECITATION"): + with self.subTest(raw=raw): + self.assertEqual(normalize_finish(raw), FinishReason.BLOCKED) + self.assertEqual(normalize_finish("TOO_MANY_TOOL_CALLS"), FinishReason.ERROR) + + def test_google_unspecified_prompt_feedback_does_not_block_valid_reply(self): + wire = self.wires(calls=False)[5][1] + wire.prompt_feedback = NS(block_reason="BLOCKED_REASON_UNSPECIFIED") + with patch("openkyrozen.providers.usage._track_cost"): + self.assertEqual(self.adapter(GoogleProvider, wire).chat_response([]).finish_reason, FinishReason.FINAL) + + def test_ollama_http_error_retains_context_overflow_detail(self): + import requests + response = requests.Response() + response.status_code = 400 + response.url = "http://localhost:11434/api/chat" + response._content = b'{"error":"context length exceeded"}' + provider = self.adapter(OllamaNativeProvider, {}) + provider._requests = NS(post=Mock(return_value=response), HTTPError=requests.HTTPError) + with self.assertRaises(requests.HTTPError) as raised: + provider.chat_response([]) + self.assertIn("context length exceeded", str(raised.exception))