From 78066fab1b4f1b2fe9134e3e276581085f87f21c Mon Sep 17 00:00:00 2001 From: Jeroen Schweitzer Date: Tue, 11 Aug 2026 17:25:18 +0200 Subject: [PATCH] style: apply ruff's automatic fixes and formatter MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Mechanical only, and separated from the judgment calls that follow so the reviewable changes are not buried in a 98-file whitespace diff. 227 automatic fixes: 60 blank lines carrying whitespace, 60 unsorted import blocks, 34 Optional[X] to X | None, 28 unused imports, 16 deprecated typing imports, 12 datetime.timezone.utc to datetime.UTC, and assorted smaller modernisations. Then `ruff format` over src and tests: 98 files reformatted, 35 already conforming. No file among the unused-import findings defines __all__ or is an __init__.py, so nothing here removes a re-export. `make test`: 658 passed, unchanged from HEAD. Two things observed while verifying, neither addressed here: `pytest tests/` cannot collect — tests/e2e/test_orchestration_e2e.py uses an `e2e` marker that is not registered, and the config is strict about markers. This fails identically at HEAD, so it predates this change; `make test` passes because it ignores tests/e2e, tests/integration and tests/contracts. test_tatlock_tool_call_logging_calculator is flaky. It failed once in a full run with these changes and passed on the next, passes in isolation with them, and fails in isolation at HEAD. It is order- or timing-dependent, not a regression from this commit — established by running the full suite both ways rather than by reasoning about which change could have caused it. Co-Authored-By: Claude --- src/agents/base.py | 12 +- src/agents/biographer/__init__.py | 1 + src/agents/biographer/agent.py | 11 +- src/agents/biographer/capability.py | 1 + src/agents/biographer/tools.py | 13 +- src/agents/delegation.py | 23 +- src/agents/housekeeper/__init__.py | 1 + src/agents/housekeeper/agent.py | 11 +- src/agents/housekeeper/capability.py | 1 + src/agents/housekeeper/client.py | 37 +-- src/agents/housekeeper/tools.py | 17 +- src/agents/librarian/__init__.py | 1 + src/agents/librarian/agent.py | 6 +- src/agents/librarian/capability.py | 1 + src/agents/librarian/client.py | 45 +++- src/agents/librarian/tools.py | 41 ++-- src/agents/lorem_tester.py | 73 ++---- src/agents/orchestration.py | 31 ++- src/agents/registry.py | 23 +- src/agents/steward/__init__.py | 1 + src/agents/steward/agent.py | 21 +- src/agents/steward/schemas.py | 27 ++- src/agents/steward/service.py | 128 ++++++---- src/agents/tatlock.py | 140 ++++++----- src/agents/tatlock_core/__init__.py | 3 +- src/agents/tatlock_core/capability.py | 2 +- src/agents/tatlock_core/tools.py | 50 ++-- src/agents/tatlock_core/toolset.py | 2 +- src/agents/tools.py | 47 ++-- src/chat/router.py | 3 +- src/chat/schemas.py | 9 + src/chat/service.py | 26 +- src/core/config.py | 120 +++------- src/core/context.py | 11 +- src/core/embeddings.py | 1 - src/core/exceptions.py | 11 +- src/core/household_registry.py | 16 +- src/core/logging_config.py | 46 +++- src/core/memory_cache.py | 3 +- src/core/memory_service.py | 28 ++- src/core/models.py | 11 +- src/core/multi_tenancy.py | 9 +- src/core/preprocessing.py | 26 +- src/core/qdrant.py | 3 +- src/core/router.py | 5 +- src/core/startup.py | 1 + src/core/tool_tracking.py | 23 +- src/core/tracing.py | 37 +-- src/core/tracing_router.py | 29 ++- src/main.py | 15 +- src/models/router.py | 5 +- src/models/schemas.py | 3 + src/ollama/client.py | 63 ++--- src/ollama/provider.py | 1 + src/ollama/schemas.py | 6 + src/responses/__init__.py | 8 +- src/responses/context.py | 21 +- src/responses/history.py | 25 +- src/responses/router.py | 15 +- src/responses/schemas.py | 82 +++---- src/responses/service.py | 131 +++++----- src/responses/streaming.py | 166 ++++++------- tests/agents/biographer/test_capability.py | 7 +- tests/agents/housekeeper/test_capability.py | 7 +- tests/agents/housekeeper/test_client.py | 3 +- tests/agents/librarian/test_capability.py | 7 +- tests/agents/librarian/test_client.py | 32 +-- tests/agents/librarian/test_client_retry.py | 20 +- .../agents/librarian/test_hybrid_contract.py | 40 +--- tests/agents/librarian/test_tool_schemas.py | 4 +- tests/agents/librarian/test_tools.py | 61 ++--- tests/agents/steward/test_steward_schemas.py | 7 +- tests/agents/steward/test_steward_service.py | 41 ++-- tests/agents/test_delegation.py | 40 +++- tests/agents/test_lorem_tester.py | 4 +- tests/agents/test_orchestration.py | 109 +++++---- tests/agents/test_protocol.py | 1 + tests/agents/test_registry.py | 10 +- tests/agents/test_tatlock_agent.py | 226 ++++++++---------- tests/agents/test_tools.py | 6 +- tests/anthropic/test_model_selector.py | 1 + tests/chat/test_router.py | 21 +- tests/chat/test_streaming_wrapper.py | 42 ++-- tests/conftest.py | 18 +- tests/contracts/test_service_contracts.py | 11 +- tests/core/test_household_registry.py | 12 +- tests/core/test_logging_config.py | 11 +- tests/core/test_memory_service.py | 5 +- tests/core/test_router.py | 7 +- tests/core/test_tenant_guard.py | 14 +- tests/core/test_tool_tracking.py | 17 +- tests/e2e/test_api_endpoints.py | 152 +++++------- tests/e2e/test_orchestration_e2e.py | 156 ++++++------ tests/integration/test_steward_streaming.py | 28 ++- .../test_steward_tatlock_integration.py | 27 ++- tests/integration/test_tatlock_streaming.py | 53 ++-- tests/models/test_router.py | 5 +- tests/responses/test_advanced_features.py | 119 ++++----- tests/responses/test_error_handling.py | 76 ++---- tests/responses/test_history.py | 95 +++----- tests/responses/test_router.py | 71 ++---- tests/responses/test_streaming_delegation.py | 22 +- tests/test_main.py | 33 ++- 103 files changed, 1601 insertions(+), 1749 deletions(-) diff --git a/src/agents/base.py b/src/agents/base.py index 3383b3f..38893c0 100644 --- a/src/agents/base.py +++ b/src/agents/base.py @@ -6,7 +6,8 @@ must implement. The interface is designed around the Responses API format. """ from abc import ABC, abstractmethod -from typing import AsyncGenerator, Any +from collections.abc import AsyncGenerator +from typing import Any class OutputItem: @@ -19,12 +20,7 @@ class OutputItem: - message: Assistant response message """ - def __init__( - self, - type: str, - id: str, - **kwargs: Any - ): + def __init__(self, type: str, id: str, **kwargs: Any): self.type = type self.id = id self.data = kwargs @@ -48,7 +44,7 @@ class AgentInterface(ABC): temperature: float = 1.0, max_tokens: int | None = None, stop: list[str] | None = None, - **kwargs: Any + **kwargs: Any, ) -> AsyncGenerator[OutputItem, None]: """ Generate streaming response as output items. diff --git a/src/agents/biographer/__init__.py b/src/agents/biographer/__init__.py index ab19c94..822b533 100644 --- a/src/agents/biographer/__init__.py +++ b/src/agents/biographer/__init__.py @@ -11,6 +11,7 @@ For direct key-based lookups (location, timezone, preferences), use the memory_service instead - it's faster and doesn't require LLM. The Biographer handles semantic, fuzzy queries. """ + from src.agents.biographer.agent import ( get_biographer_agent, run_biographer, diff --git a/src/agents/biographer/agent.py b/src/agents/biographer/agent.py index 71a4eb4..6fe7e7c 100644 --- a/src/agents/biographer/agent.py +++ b/src/agents/biographer/agent.py @@ -7,7 +7,8 @@ A PydanticAI agent that serves as the household's memory keeper: - Manages user profile and preferences - Forgets information when requested """ -from typing import Any, Optional + +from typing import Any from pydantic_ai import Agent @@ -19,7 +20,6 @@ from src.agents.biographer.tools import ( update_preference, update_profile, ) -from src.core.config import config from src.core.logging_config import get_logger logger = get_logger(__name__) @@ -97,7 +97,7 @@ When recalling: """ # Lazy initialization to avoid connection issues during imports -_biographer_agent: Optional[Agent[None, str]] = None +_biographer_agent: Agent[None, str] | None = None def _create_biographer_agent() -> Agent[None, str]: @@ -126,6 +126,7 @@ def _create_biographer_agent() -> Agent[None, str]: agent.tool_plain(forget_memory) from src.anthropic.model_selector import get_model_info + model_info = get_model_info() logger.info( "biographer_agent_created", @@ -153,7 +154,7 @@ def get_biographer_agent() -> Agent[None, str]: async def run_biographer( task: str, context: str = "", - message_history: Optional[list[Any]] = None, + message_history: list[Any] | None = None, ) -> str: """ Execute a memory task with The Biographer. @@ -216,7 +217,7 @@ async def run_biographer( async def run_biographer_stream( task: str, context: str = "", - message_history: Optional[list[Any]] = None, + message_history: list[Any] | None = None, ): """ Execute a memory task with streaming output. diff --git a/src/agents/biographer/capability.py b/src/agents/biographer/capability.py index 913e24e..d972595 100644 --- a/src/agents/biographer/capability.py +++ b/src/agents/biographer/capability.py @@ -4,6 +4,7 @@ Biographer capability registration for the Household Registry. Defines The Biographer's capabilities and registers it as a household member for coordination by the Steward and Tatlock. """ + from src.agents.biographer.agent import get_biographer_agent from src.agents.biographer.tools import BIOGRAPHER_TOOLS from src.core.household_registry import ( diff --git a/src/agents/biographer/tools.py b/src/agents/biographer/tools.py index f9f920a..642555d 100644 --- a/src/agents/biographer/tools.py +++ b/src/agents/biographer/tools.py @@ -10,6 +10,7 @@ These tools enable The Biographer to record and recall the user's story: For direct key-based access (get/set profile, preferences), use memory_service directly - these tools are for semantic queries. """ + from src.core.context import get_user from src.core.embeddings import get_embedding_client from src.core.logging_config import get_logger @@ -23,6 +24,7 @@ logger = get_logger(__name__) # Semantic Recall # ============================================================================ + async def recall_semantic( query: str, memory_type: str = "", @@ -108,6 +110,7 @@ async def recall_semantic( # Store Memory # ============================================================================ + async def store_insight( key: str, value: str, @@ -158,7 +161,7 @@ async def store_insight( f"**Keywords:** {', '.join(keywords)}", f"**Importance:** {importance:.1f}", "", - "_Memory is now searchable via semantic recall._" + "_Memory is now searchable via semantic recall._", ] logger.info( @@ -216,7 +219,7 @@ async def update_profile( "## Profile Updated", f"**{key}:** {value}", "", - "_Profile data is automatically included in context._" + "_Profile data is automatically included in context._", ] logger.info( @@ -271,7 +274,7 @@ async def update_preference( "## Preference Updated", f"**{key}:** {value}", "", - "_Preference will be applied to future responses._" + "_Preference will be applied to future responses._", ] logger.info( @@ -293,6 +296,7 @@ async def update_preference( # List Memories # ============================================================================ + async def list_memories( memory_type: str = "learned_fact", limit: int = 20, @@ -378,6 +382,7 @@ async def list_memories( # Forget Memory # ============================================================================ + async def forget_memory( key: str, memory_type: str = "learned_fact", @@ -420,7 +425,7 @@ async def forget_memory( f"**Key:** {key}", f"**Type:** {memory_type}", "", - "_Memory has been removed._" + "_Memory has been removed._", ] logger.info( diff --git a/src/agents/delegation.py b/src/agents/delegation.py index 8b3ef2a..b04c356 100644 --- a/src/agents/delegation.py +++ b/src/agents/delegation.py @@ -8,6 +8,7 @@ returns a structured result for synthesis. This implements the agent-as-tool pattern recommended by PydanticAI: agents call other agents via tool wrappers, keeping each agent focused. """ + import asyncio from dataclasses import dataclass, field from enum import Enum @@ -23,6 +24,7 @@ logger = get_logger(__name__) # Action Types for Think Slug Selection # ============================================================================= + class ActionType(Enum): """ Categories of actions for selecting appropriate think messages. @@ -30,11 +32,12 @@ class ActionType(Enum): Each expert has different action types that warrant different butler-perspective messages to the user. """ - RETRIEVE = "retrieve" # Looking up existing information - RESEARCH = "research" # Conducting new research (web search, etc.) - CREATE = "create" # Creating new content (pages, notes) - CONTROL = "control" # Controlling devices/automations - RECORD = "record" # Recording memories/notes + + RETRIEVE = "retrieve" # Looking up existing information + RESEARCH = "research" # Conducting new research (web search, etc.) + CREATE = "create" # Creating new content (pages, notes) + CONTROL = "control" # Controlling devices/automations + RECORD = "record" # Recording memories/notes # ============================================================================= @@ -159,8 +162,7 @@ def build_delegation_context( if isinstance(content, list): # Tolerate structured content parts content = " ".join( - part.get("text", "") if isinstance(part, dict) else str(part) - for part in content + part.get("text", "") if isinstance(part, dict) else str(part) for part in content ) content = str(content).strip() if content: @@ -206,6 +208,7 @@ class DelegationTask: depends_on: List of task IDs this task depends on result: Result from expert after execution """ + expert_name: str task: str context: str = "" @@ -219,6 +222,7 @@ class DelegationTask: """Generate task ID if not provided.""" if not self.task_id: import uuid + self.task_id = f"{self.expert_name}_{uuid.uuid4().hex[:8]}" @@ -236,6 +240,7 @@ class DelegationResult: error: Short user-safe error label if failed. Exception detail stays in the logs only """ + expert_name: str task: str success: bool @@ -335,9 +340,7 @@ async def delegate_to_librarian( if span: span.metadata["success"] = False - span.details["error"] = ( - f"timed out after {config.LIBRARIAN_TIMEOUT}s" - ) + span.details["error"] = f"timed out after {config.LIBRARIAN_TIMEOUT}s" return DelegationResult( expert_name="librarian", diff --git a/src/agents/housekeeper/__init__.py b/src/agents/housekeeper/__init__.py index ed52852..ea1e880 100644 --- a/src/agents/housekeeper/__init__.py +++ b/src/agents/housekeeper/__init__.py @@ -4,6 +4,7 @@ The Housekeeper - Home Automation Agent. Provides home automation capabilities through the core-api service, which wraps the Home Assistant REST API into LLM-friendly endpoints. """ + from src.agents.housekeeper.agent import run_housekeeper, run_housekeeper_stream from src.agents.housekeeper.capability import ( HOUSEKEEPER_CAPABILITY, diff --git a/src/agents/housekeeper/agent.py b/src/agents/housekeeper/agent.py index df30824..d509e9b 100644 --- a/src/agents/housekeeper/agent.py +++ b/src/agents/housekeeper/agent.py @@ -8,7 +8,8 @@ the core-api service, which wraps Home Assistant REST API, offering: - Script execution - Automation management """ -from typing import Any, Optional + +from typing import Any from pydantic_ai import Agent @@ -27,7 +28,6 @@ from src.agents.housekeeper.tools import ( turn_off, turn_on, ) -from src.core.config import config from src.core.logging_config import get_logger logger = get_logger(__name__) @@ -98,7 +98,7 @@ After completing actions, briefly confirm: """ # Lazy initialization to avoid connection issues during imports -_housekeeper_agent: Optional[Agent[None, str]] = None +_housekeeper_agent: Agent[None, str] | None = None def _create_housekeeper_agent() -> Agent[None, str]: @@ -140,6 +140,7 @@ def _create_housekeeper_agent() -> Agent[None, str]: agent.tool_plain(get_history) from src.anthropic.model_selector import get_model_info + model_info = get_model_info() logger.info( "housekeeper_agent_created", @@ -167,7 +168,7 @@ def get_housekeeper_agent() -> Agent[None, str]: async def run_housekeeper( task: str, context: str = "", - message_history: Optional[list[Any]] = None, + message_history: list[Any] | None = None, ) -> str: """ Execute a home automation task with The Housekeeper. @@ -234,7 +235,7 @@ async def run_housekeeper( async def run_housekeeper_stream( task: str, context: str = "", - message_history: Optional[list[Any]] = None, + message_history: list[Any] | None = None, ): """ Execute a home automation task with streaming output. diff --git a/src/agents/housekeeper/capability.py b/src/agents/housekeeper/capability.py index bf2c8fc..9a83061 100644 --- a/src/agents/housekeeper/capability.py +++ b/src/agents/housekeeper/capability.py @@ -4,6 +4,7 @@ Housekeeper capability registration for the Household Registry. Defines The Housekeeper's capabilities and registers it as a household member for coordination by the Steward and Tatlock. """ + from src.agents.housekeeper.agent import get_housekeeper_agent from src.agents.housekeeper.tools import HOUSEKEEPER_TOOLS from src.core.household_registry import ( diff --git a/src/agents/housekeeper/client.py b/src/agents/housekeeper/client.py index db50f97..3d65c2c 100644 --- a/src/agents/housekeeper/client.py +++ b/src/agents/housekeeper/client.py @@ -5,7 +5,8 @@ Provides async methods for home automation operations via Home Assistant. Core-API is a separate service that wraps the Home Assistant REST API into LLM-friendly endpoints. """ -from typing import Any, Optional + +from typing import Any import httpx from pydantic import BaseModel, Field @@ -28,7 +29,7 @@ class Device(BaseModel): name: str state: str domain: str - area: Optional[str] = None + area: str | None = None attributes: dict[str, Any] = Field(default_factory=dict) @@ -38,8 +39,8 @@ class DeviceState(BaseModel): entity_id: str state: str attributes: dict[str, Any] = Field(default_factory=dict) - last_changed: Optional[str] = None - last_updated: Optional[str] = None + last_changed: str | None = None + last_updated: str | None = None class Scene(BaseModel): @@ -47,7 +48,7 @@ class Scene(BaseModel): entity_id: str name: str - friendly_name: Optional[str] = None + friendly_name: str | None = None class Script(BaseModel): @@ -55,8 +56,8 @@ class Script(BaseModel): entity_id: str name: str - description: Optional[str] = None - last_triggered: Optional[str] = None + description: str | None = None + last_triggered: str | None = None class Automation(BaseModel): @@ -64,8 +65,8 @@ class Automation(BaseModel): entity_id: str name: str - state: str = "on" - last_triggered: Optional[str] = None + state: str = "on" + last_triggered: str | None = None class HistoryEntry(BaseModel): @@ -109,8 +110,8 @@ class CoreAPIClient: def __init__( self, - base_url: Optional[str] = None, - api_key: Optional[str] = None, + base_url: str | None = None, + api_key: str | None = None, timeout: int = 30, ): """ @@ -124,7 +125,7 @@ class CoreAPIClient: self.base_url = base_url or str(config.CORE_API_HOST) self.api_key = api_key or config.CORE_API_KEY self.timeout = timeout - self._client: Optional[httpx.AsyncClient] = None + self._client: httpx.AsyncClient | None = None async def __aenter__(self) -> "CoreAPIClient": """Create HTTP client on context entry.""" @@ -159,8 +160,8 @@ class CoreAPIClient: async def list_devices( self, - domain: Optional[str] = None, - area: Optional[str] = None, + domain: str | None = None, + area: str | None = None, ) -> list[Device]: """ List devices, optionally filtered by domain or area. @@ -231,9 +232,9 @@ class CoreAPIClient: async def turn_on( self, entity_id: str, - brightness: Optional[int] = None, - color_temp: Optional[int] = None, - rgb_color: Optional[tuple[int, int, int]] = None, + brightness: int | None = None, + color_temp: int | None = None, + rgb_color: tuple[int, int, int] | None = None, ) -> ControlResult: """ Turn on a device. @@ -399,7 +400,7 @@ class CoreAPIClient: async def run_script( self, script_id: str, - variables: Optional[dict[str, Any]] = None, + variables: dict[str, Any] | None = None, ) -> ControlResult: """ Run a script. diff --git a/src/agents/housekeeper/tools.py b/src/agents/housekeeper/tools.py index e95946d..77dcf4b 100644 --- a/src/agents/housekeeper/tools.py +++ b/src/agents/housekeeper/tools.py @@ -4,6 +4,7 @@ Housekeeper tools for PydanticAI agent. These tools wrap the core-api service and are registered with The Housekeeper agent for home automation tasks. """ + from src.agents.housekeeper.client import CoreAPIClient from src.core.logging_config import get_logger @@ -72,14 +73,24 @@ async def list_devices( return True return False - sorted_devices = sorted(dom_devices, key=lambda d: (not is_room_group(d), d.entity_id)) + sorted_devices = sorted( + dom_devices, key=lambda d: (not is_room_group(d), d.entity_id) + ) for device in sorted_devices: - state_icon = "on" if device.state == "on" else "off" if device.state == "off" else device.state + state_icon = ( + "on" + if device.state == "on" + else "off" + if device.state == "off" + else device.state + ) area_str = f" ({device.area})" if device.area else "" # Mark room groups clearly using actual HA data group_marker = " [ROOM GROUP]" if is_room_group(device) else "" - output_parts.append(f"- **{device.name}**{area_str}{group_marker}: {state_icon}") + output_parts.append( + f"- **{device.name}**{area_str}{group_marker}: {state_icon}" + ) output_parts.append(f" ID: `{device.entity_id}`") output_parts.append("") diff --git a/src/agents/librarian/__init__.py b/src/agents/librarian/__init__.py index d5f18d0..63deeff 100644 --- a/src/agents/librarian/__init__.py +++ b/src/agents/librarian/__init__.py @@ -7,6 +7,7 @@ Connects to the library-desk API to provide: - Knowledge graph queries - Semantic search """ + from src.agents.librarian.agent import ( get_librarian_agent, run_librarian, diff --git a/src/agents/librarian/agent.py b/src/agents/librarian/agent.py index 32318e5..65f7a9a 100644 --- a/src/agents/librarian/agent.py +++ b/src/agents/librarian/agent.py @@ -7,6 +7,7 @@ the library-desk API, offering: - Wiki and document management - Semantic search and knowledge graph exploration """ + from typing import Any from pydantic_ai import Agent @@ -202,6 +203,7 @@ def _create_librarian_agent() -> Agent[None, str]: agent.tool_plain(smart_create_wiki_page) from src.anthropic.model_selector import get_model_info + model_info = get_model_info() logger.info( "librarian_agent_created", @@ -294,6 +296,4 @@ async def run_librarian( error=str(e), exc_info=True, ) - raise AgentError( - "Research task failed", agent_name="librarian" - ) from e + raise AgentError("Research task failed", agent_name="librarian") from e diff --git a/src/agents/librarian/capability.py b/src/agents/librarian/capability.py index 24ab07b..ac8a540 100644 --- a/src/agents/librarian/capability.py +++ b/src/agents/librarian/capability.py @@ -4,6 +4,7 @@ Librarian capability registration for the Household Registry. Defines The Librarian's capabilities and registers it as a household member for coordination by the Steward and Tatlock. """ + from src.agents.librarian.agent import get_librarian_agent from src.agents.librarian.tools import LIBRARIAN_TOOLS from src.core.household_registry import ( diff --git a/src/agents/librarian/client.py b/src/agents/librarian/client.py index 7ae6b53..962793b 100644 --- a/src/agents/librarian/client.py +++ b/src/agents/librarian/client.py @@ -7,6 +7,7 @@ Provides async methods for all relevant library-desk endpoints: - Vector search - Knowledge graph queries """ + import asyncio from collections.abc import AsyncIterator, Awaitable, Callable from contextlib import asynccontextmanager @@ -38,8 +39,10 @@ _shared_http_client: ContextVar[httpx.AsyncClient | None] = ContextVar( # Response Models # ============================================================================ + class WikiPage(BaseModel): """Wiki page from library-desk.""" + id: int path: str title: str @@ -52,6 +55,7 @@ class WikiPage(BaseModel): class WikiSearchResult(BaseModel): """Search result from wiki search.""" + id: int path: str title: str @@ -61,6 +65,7 @@ class WikiSearchResult(BaseModel): class VectorSearchResult(BaseModel): """Result from semantic vector search.""" + page_id: int page_path: str page_title: str @@ -71,8 +76,11 @@ class VectorSearchResult(BaseModel): class HybridSearchResult(BaseModel): """Result from HybridRAG search.""" + source: str # source_type: "wiki", "web", "volatile", "document" - sources: list[str] = Field(default_factory=list) # legs that found it: "vector", "graph", "web", ... + sources: list[str] = Field( + default_factory=list + ) # legs that found it: "vector", "graph", "web", ... title: str content: str url: str | None = None @@ -84,6 +92,7 @@ class HybridSearchResult(BaseModel): class HybridRAGResponse(BaseModel): """Full response from HybridRAG query.""" + results: list[HybridSearchResult] = Field(default_factory=list) keywords: list[str] = Field(default_factory=list) synonyms: list[str] = Field(default_factory=list) @@ -102,6 +111,7 @@ class HybridRAGResponse(BaseModel): class GraphNode(BaseModel): """Node from knowledge graph.""" + id: str labels: list[str] = Field(default_factory=list) properties: dict[str, Any] = Field(default_factory=dict) @@ -109,12 +119,14 @@ class GraphNode(BaseModel): class Dossier(BaseModel): """A dossier (tag-based collection).""" + name: str page_count: int class ResearchSummary(BaseModel): """Summary of research performed during smart-create.""" + wiki_results: int = 0 web_results: int = 0 graph_entities: int = 0 @@ -124,6 +136,7 @@ class ResearchSummary(BaseModel): class WebSearchResult(BaseModel): """Result from web search via /rag/search.""" + title: str url: str content: str = "" # Full extracted text via Trafilatura @@ -134,6 +147,7 @@ class WebSearchResult(BaseModel): class WebSearchResponse(BaseModel): """Response from /rag/search endpoint.""" + query: str search_type: str results: list[WebSearchResult] = Field(default_factory=list) @@ -144,6 +158,7 @@ class WebSearchResponse(BaseModel): class ContentExtractionResult(BaseModel): """Result from content extraction.""" + url: str title: str | None = None content: str = "" @@ -156,6 +171,7 @@ class ContentExtractionResult(BaseModel): class BatchExtractionResponse(BaseModel): """Response from batch content extraction.""" + results: list[ContentExtractionResult] = Field(default_factory=list) total_urls: int = 0 successful: int = 0 @@ -165,6 +181,7 @@ class BatchExtractionResponse(BaseModel): class EntityLinking(BaseModel): """Entity linking results from smart-create.""" + forward_links: int = 0 backward_links: int = 0 pages_updated: int = 0 @@ -172,6 +189,7 @@ class EntityLinking(BaseModel): class SmartCreateResponse(BaseModel): """Response from smart-create wiki page endpoint.""" + page: WikiPage research_summary: ResearchSummary = Field(default_factory=ResearchSummary) sources_used: int = 0 @@ -183,6 +201,7 @@ class SmartCreateResponse(BaseModel): # Client # ============================================================================ + class LibraryDeskClient: """ Async HTTP client for Library-Desk API. @@ -398,17 +417,19 @@ class LibraryDeskClient: # related_dossiers; older names kept as fallbacks) results = [] for r in data.get("results", []): - results.append(HybridSearchResult( - source=r.get("source_type") or r.get("source", "unknown"), - sources=r.get("sources", []), - title=r.get("title", ""), - content=r.get("content", ""), - url=r.get("url"), - score=r.get("rrf_score", r.get("score", 0.0)), - page_id=r.get("page_id"), - related_dossiers=r.get("related_dossiers", []), - metadata=r.get("metadata", {}), - )) + results.append( + HybridSearchResult( + source=r.get("source_type") or r.get("source", "unknown"), + sources=r.get("sources", []), + title=r.get("title", ""), + content=r.get("content", ""), + url=r.get("url"), + score=r.get("rrf_score", r.get("score", 0.0)), + page_id=r.get("page_id"), + related_dossiers=r.get("related_dossiers", []), + metadata=r.get("metadata", {}), + ) + ) # Handle keywords being either a list or a dict with core_keywords; # the live service nests synonyms inside the keywords dict as a diff --git a/src/agents/librarian/tools.py b/src/agents/librarian/tools.py index e6669fe..9b1868a 100644 --- a/src/agents/librarian/tools.py +++ b/src/agents/librarian/tools.py @@ -4,6 +4,7 @@ Librarian tools for PydanticAI agent. These tools wrap the library-desk API and are registered with The Librarian agent for research and knowledge management tasks. """ + import httpx from pydantic_ai import ModelRetry @@ -27,9 +28,8 @@ def _retry_if_transient(e: Exception, what: str) -> None: status = e.response.status_code retryable = status >= 500 or status == 429 if retryable: - raise ModelRetry( - f"{what} is temporarily unavailable; please retry." - ) from e + raise ModelRetry(f"{what} is temporarily unavailable; please retry.") from e + # Icons keyed by the values library-desk emits in each result's `sources` # list (search legs) and `source_type` (result origin). @@ -64,11 +64,7 @@ def _coverage_note( results, so their absence is normal ranking behavior, not an outage. """ if response.source_status: - failed = sorted( - leg - for leg, status in response.source_status.items() - if status == "failed" - ) + failed = sorted(leg for leg, status in response.source_status.items() if status == "failed") if failed: return ( "⚠️ *Coverage note: results are partial - " @@ -111,6 +107,7 @@ def _coverage_note( # HybridRAG Search # ============================================================================ + async def hybrid_search( query: str, include_web: bool = True, @@ -165,9 +162,7 @@ async def hybrid_search( # Add related dossiers if response.related_dossiers: - output_parts.append( - f"**Related Dossiers:** {', '.join(response.related_dossiers)}" - ) + output_parts.append(f"**Related Dossiers:** {', '.join(response.related_dossiers)}") output_parts.append("") @@ -216,6 +211,7 @@ async def hybrid_search( # Wiki Operations # ============================================================================ + async def search_wiki( query: str, limit: int = 10, @@ -331,9 +327,7 @@ async def list_dossiers() -> str: output_parts = ["## Research Dossiers\n"] for dossier in dossiers: - output_parts.append( - f"- **{dossier.name}** ({dossier.page_count} pages)" - ) + output_parts.append(f"- **{dossier.name}** ({dossier.page_count} pages)") return "\n".join(output_parts) @@ -389,6 +383,7 @@ async def get_dossier_pages( # Semantic Search # ============================================================================ + async def semantic_search( query: str, limit: int = 10, @@ -420,9 +415,7 @@ async def semantic_search( output_parts = [f"## Semantic Search: {query}\n"] for i, result in enumerate(results, 1): - output_parts.append( - f"{i}. **{result.page_title}** (score: {result.score:.2f})" - ) + output_parts.append(f"{i}. **{result.page_title}** (score: {result.score:.2f})") output_parts.append(f" Path: {result.page_path}") output_parts.append(f" {result.chunk_text[:200]}...") output_parts.append("") @@ -439,6 +432,7 @@ async def semantic_search( # Knowledge Graph # ============================================================================ + async def explore_knowledge_graph( entity_type: str = "Document", limit: int = 20, @@ -565,6 +559,7 @@ async def find_related_entities( # Web Search & Content Extraction # ============================================================================ + async def search_web( query: str, limit: int = 10, @@ -605,7 +600,9 @@ async def search_web( return f"No results found for '{query}'" output_parts = [f"## Web Search: {query}\n"] - output_parts.append(f"*Found {response.total_results} results in {response.search_time_ms}ms*\n") + output_parts.append( + f"*Found {response.total_results} results in {response.search_time_ms}ms*\n" + ) for i, result in enumerate(response.results, 1): output_parts.append(f"### {i}. {result.title}") @@ -880,7 +877,9 @@ async def update_wiki_page( if page.tags: output_parts.append(f"**Tags:** {', '.join(page.tags)}") - output_parts.append("\n*Vector embeddings and knowledge graph will be updated automatically.*") + output_parts.append( + "\n*Vector embeddings and knowledge graph will be updated automatically.*" + ) logger.info( "librarian_update_page", @@ -954,7 +953,9 @@ async def create_wiki_page( if page.description: output_parts.append(f"**Description:** {page.description}") - output_parts.append("\n*Vector embeddings and knowledge graph will be updated automatically.*") + output_parts.append( + "\n*Vector embeddings and knowledge graph will be updated automatically.*" + ) logger.info( "librarian_create_page", diff --git a/src/agents/lorem_tester.py b/src/agents/lorem_tester.py index e2c74f6..70d3be8 100644 --- a/src/agents/lorem_tester.py +++ b/src/agents/lorem_tester.py @@ -12,16 +12,16 @@ infrastructure is real production code. import asyncio import random import secrets -from typing import AsyncGenerator, Any +from collections.abc import AsyncGenerator +from typing import Any from src.agents.base import AgentInterface, OutputItem from src.core.exceptions import ( - RateLimitError, - ContextLengthError, APIError, + ContextLengthError, + RateLimitError, ) - # Mock lorem ipsum content LOREM_PARAGRAPHS = [ "Lorem ipsum dolor sit amet, consectetur adipiscing elit. Sed do eiusmod tempor incididunt ut labore et dolore magna aliqua.", @@ -46,33 +46,27 @@ MOCK_TOOLS = [ "description": "Search the knowledge base for relevant information", "parameters": { "type": "object", - "properties": { - "query": {"type": "string", "description": "Search query"} - }, - "required": ["query"] - } + "properties": {"query": {"type": "string", "description": "Search query"}}, + "required": ["query"], + }, }, { "name": "calculate", "description": "Perform mathematical calculations", "parameters": { "type": "object", - "properties": { - "expression": {"type": "string", "description": "Math expression"} - }, - "required": ["expression"] - } + "properties": {"expression": {"type": "string", "description": "Math expression"}}, + "required": ["expression"], + }, }, { "name": "get_weather", "description": "Get current weather for a location", "parameters": { "type": "object", - "properties": { - "location": {"type": "string", "description": "City name"} - }, - "required": ["location"] - } + "properties": {"location": {"type": "string", "description": "City name"}}, + "required": ["location"], + }, }, ] @@ -109,7 +103,7 @@ class LoremTesterAgent(AgentInterface): temperature: float = 1.0, max_tokens: int | None = None, stop: list[str] | None = None, - **kwargs: Any + **kwargs: Any, ) -> AsyncGenerator[OutputItem, None]: """ Generate mock response with reasoning, tools, and content. @@ -123,8 +117,7 @@ class LoremTesterAgent(AgentInterface): # 1. Yield reasoning item if requested if reasoning and reasoning.get("summary") == "auto": yield await self._create_reasoning_item( - messages, - effort=reasoning.get("effort", "medium") + messages, effort=reasoning.get("effort", "medium") ) # 2. Randomly yield function calls if tools available (30% chance) @@ -150,7 +143,7 @@ class LoremTesterAgent(AgentInterface): "reasoning": True, "tools": True, "vision": False, # Not yet - "audio": False, # Not yet + "audio": False, # Not yet } # Private helper methods @@ -179,9 +172,7 @@ class LoremTesterAgent(AgentInterface): raise APIError("Invalid tool call: tool 'nonexistent' not found (mock trigger)") async def _create_reasoning_item( - self, - messages: list[dict], - effort: str = "medium" + self, messages: list[dict], effort: str = "medium" ) -> OutputItem: """Create a reasoning output item with mock thinking steps.""" @@ -200,23 +191,17 @@ class LoremTesterAgent(AgentInterface): steps = random.sample(REASONING_STEPS, min(num_steps, len(REASONING_STEPS))) return OutputItem( - type="reasoning", - id=f"rs_{generate_id()}", - summary=steps, - status="completed" + type="reasoning", id=f"rs_{generate_id()}", summary=steps, status="completed" ) - async def _create_tool_calls( - self, - tools: list[dict] - ) -> AsyncGenerator[OutputItem, None]: + async def _create_tool_calls(self, tools: list[dict]) -> AsyncGenerator[OutputItem, None]: """Create mock function call output items.""" # Randomly select 1-2 tools to "call" num_calls = random.randint(1, 2) selected_tools = random.sample( - MOCK_TOOLS[:min(len(MOCK_TOOLS), len(tools))], - min(num_calls, len(MOCK_TOOLS), len(tools)) + MOCK_TOOLS[: min(len(MOCK_TOOLS), len(tools))], + min(num_calls, len(MOCK_TOOLS), len(tools)), ) for tool in selected_tools: @@ -228,7 +213,7 @@ class LoremTesterAgent(AgentInterface): id=f"fc_{generate_id()}", name=tool["name"], arguments=args, - status="completed" + status="completed", ) def _generate_mock_args(self, tool: dict) -> str: @@ -254,11 +239,7 @@ class LoremTesterAgent(AgentInterface): # Generic mock arguments return json.dumps({"input": "mock_value"}) - async def _create_message_item( - self, - messages: list[dict], - temperature: float - ) -> OutputItem: + async def _create_message_item(self, messages: list[dict], temperature: float) -> OutputItem: """Create final message output item with lorem ipsum content.""" # Select random lorem ipsum paragraphs @@ -270,10 +251,6 @@ class LoremTesterAgent(AgentInterface): type="message", id=f"msg_{generate_id()}", role="assistant", - content=[{ - "type": "output_text", - "text": content, - "annotations": [] - }], - status="completed" + content=[{"type": "output_text", "text": content, "annotations": []}], + status="completed", ) diff --git a/src/agents/orchestration.py b/src/agents/orchestration.py index abe728f..40100a9 100644 --- a/src/agents/orchestration.py +++ b/src/agents/orchestration.py @@ -14,12 +14,13 @@ Supports: - Result aggregation from multiple experts - Partial failure handling """ + import asyncio +from collections.abc import AsyncGenerator from dataclasses import dataclass, field from enum import Enum -from typing import AsyncGenerator, Optional, Callable, Any -from src.agents.delegation import DelegationTask, DelegationResult, delegate_to_librarian +from src.agents.delegation import DelegationResult, DelegationTask, delegate_to_librarian from src.core.logging_config import get_logger logger = get_logger(__name__) @@ -27,8 +28,9 @@ logger = get_logger(__name__) class ExecutionMode(str, Enum): """Execution mode for multi-expert coordination.""" + SEQUENTIAL = "sequential" # One at a time, in order - PARALLEL = "parallel" # All at once, concurrently + PARALLEL = "parallel" # All at once, concurrently @dataclass @@ -38,12 +40,13 @@ class OrchestrationContext: Tracks the user's request, delegation tasks, and results. """ + user_message: str steward_note: str - conversation_id: Optional[str] = None + conversation_id: str | None = None -def parse_delegation_from_steward_note(steward_note: str) -> Optional[DelegationTask]: +def parse_delegation_from_steward_note(steward_note: str) -> DelegationTask | None: """ Parse a delegation task from Steward's note. @@ -68,9 +71,9 @@ def parse_delegation_from_steward_note(steward_note: str) -> Optional[Delegation # Look for DELEGATE: pattern # Match: "DELEGATE: expert_name to action description" match = re.search( - r'DELEGATE:\s*(\w+)\s+to\s+(.+?)(?:\n|REASON:|COMPLEXITY:|CONTEXT:|$)', + r"DELEGATE:\s*(\w+)\s+to\s+(.+?)(?:\n|REASON:|COMPLEXITY:|CONTEXT:|$)", steward_note, - re.IGNORECASE | re.MULTILINE + re.IGNORECASE | re.MULTILINE, ) if match: @@ -135,7 +138,7 @@ async def execute_delegation( async def orchestrate_with_think_updates( user_message: str, steward_note: str, - delegation_task: Optional[DelegationTask] = None, + delegation_task: DelegationTask | None = None, ) -> AsyncGenerator[str, None]: """ Orchestrate expert delegation with streaming think updates. @@ -218,17 +221,21 @@ def extract_delegation_context( } # Extract REASON: - reason_match = re.search(r'REASON:\s*(.+?)(?:\n|COMPLEXITY:|CONTEXT:|$)', steward_note, re.IGNORECASE) + reason_match = re.search( + r"REASON:\s*(.+?)(?:\n|COMPLEXITY:|CONTEXT:|$)", steward_note, re.IGNORECASE + ) if reason_match: result["reason"] = reason_match.group(1).strip() # Extract COMPLEXITY: - complexity_match = re.search(r'COMPLEXITY:\s*(.+?)(?:\n|CONTEXT:|$)', steward_note, re.IGNORECASE) + complexity_match = re.search( + r"COMPLEXITY:\s*(.+?)(?:\n|CONTEXT:|$)", steward_note, re.IGNORECASE + ) if complexity_match: result["complexity"] = complexity_match.group(1).strip() # Extract CONTEXT: - context_match = re.search(r'CONTEXT:\s*(.+?)$', steward_note, re.IGNORECASE | re.MULTILINE) + context_match = re.search(r"CONTEXT:\s*(.+?)$", steward_note, re.IGNORECASE | re.MULTILINE) if context_match: result["context"] = context_match.group(1).strip() @@ -239,6 +246,7 @@ def extract_delegation_context( # Multi-Expert Coordination # ============================================================================ + @dataclass class MultiExpertResult: """ @@ -250,6 +258,7 @@ class MultiExpertResult: failed_experts: List of expert names that failed combined_output: Aggregated output from all successful experts """ + results: dict[str, DelegationResult] = field(default_factory=dict) all_succeeded: bool = True failed_experts: list[str] = field(default_factory=list) diff --git a/src/agents/registry.py b/src/agents/registry.py index 51678f9..30552f8 100644 --- a/src/agents/registry.py +++ b/src/agents/registry.py @@ -8,9 +8,6 @@ It provides a central place to: - Check model capabilities """ -import time -from typing import Type - from src.agents.base import AgentInterface from src.agents.lorem_tester import LoremTesterAgent from src.agents.tatlock import TatlockAgent @@ -63,7 +60,7 @@ class ModelRegistry: if model_id not in cls.MODELS: raise ModelNotFoundError(model_id) - agent_class: Type[AgentInterface] = cls.MODELS[model_id]["agent_class"] + agent_class: type[AgentInterface] = cls.MODELS[model_id]["agent_class"] return agent_class() @classmethod @@ -106,14 +103,16 @@ class ModelRegistry: agent = cls.get_agent(model_id) capabilities = await agent.get_capabilities() - models.append({ - "id": model_id, - "object": "model", - "created": config["created"], - "owned_by": config["owned_by"], - "capabilities": capabilities, - "description": config["description"], - }) + models.append( + { + "id": model_id, + "object": "model", + "created": config["created"], + "owned_by": config["owned_by"], + "capabilities": capabilities, + "description": config["description"], + } + ) return models diff --git a/src/agents/steward/__init__.py b/src/agents/steward/__init__.py index d4d328b..2f290eb 100644 --- a/src/agents/steward/__init__.py +++ b/src/agents/steward/__init__.py @@ -4,6 +4,7 @@ Steward agent package. The Steward analyzes incoming requests and recommends relevant household capabilities, creating a two-tier architecture with the Butler. """ + from .agent import StewardAgent, get_steward_agent from .schemas import ConversationContext, StewardRecommendation from .service import analyze_request, format_steward_note diff --git a/src/agents/steward/agent.py b/src/agents/steward/agent.py index 6f2f6ef..75e2b7d 100644 --- a/src/agents/steward/agent.py +++ b/src/agents/steward/agent.py @@ -8,8 +8,8 @@ This creates a two-tier architecture that prevents cognitive overload. Uses plain text output (not JSON) for reliability. Supports both Claude (preferred) and Ollama (fallback) backends via direct API calls. """ + import httpx -from typing import Optional from src.anthropic.model_selector import get_model_info, is_claude_available, resolve_backend from src.core.config import config @@ -29,9 +29,7 @@ def build_steward_prompt(query: str, conversation_history: list[dict]) -> str: cap_list = [] for cap in capabilities: - cap_list.append( - f"• {cap.name} - {cap.description} (domains: {', '.join(cap.domains)})" - ) + cap_list.append(f"• {cap.name} - {cap.description} (domains: {', '.join(cap.domains)})") capabilities_text = "\n".join(cap_list) # Format conversation history if present @@ -114,7 +112,7 @@ class StewardAgent: def __init__(self): """Initialize Steward with backend selection based on availability.""" # Ollama config (primary) - self.ollama_host = str(config.OLLAMA_HOST).rstrip('/') + self.ollama_host = str(config.OLLAMA_HOST).rstrip("/") self.ollama_model = config.OLLAMA_DEFAULT_MODEL # Claude config (fallback) @@ -139,6 +137,7 @@ class StewardAgent: """Get or create Anthropic client (lazy initialization).""" if self._anthropic_client is None: from anthropic import AsyncAnthropic + self._anthropic_client = AsyncAnthropic(api_key=config.ANTHROPIC_API_KEY) return self._anthropic_client @@ -167,20 +166,16 @@ class StewardAgent: "stream": False, "options": { "temperature": 0.3, # Lower = more consistent - "top_p": 0.9 - } - } + "top_p": 0.9, + }, + }, ) response.raise_for_status() result = response.json() return result["response"].strip() - async def analyze( - self, - query: str, - conversation_history: Optional[list[dict]] = None - ) -> str: + async def analyze(self, query: str, conversation_history: list[dict] | None = None) -> str: """ Analyze query and return plain text recommendation. diff --git a/src/agents/steward/schemas.py b/src/agents/steward/schemas.py index 377ce58..2b37b3c 100644 --- a/src/agents/steward/schemas.py +++ b/src/agents/steward/schemas.py @@ -4,7 +4,8 @@ Steward agent schemas. Defines the structured output models for Steward's request analysis and capability recommendations. """ -from typing import Any, Literal, Optional + +from typing import Any, Literal from pydantic import BaseModel, Field @@ -16,16 +17,16 @@ class ConversationContext(BaseModel): The Steward analyzes the full conversation to identify references to previous topics, helping the Butler maintain context. """ + has_previous_context: bool = Field( description="Whether the current request references previous conversation turns" ) relevant_turns: list[int] = Field( default_factory=list, - description="0-indexed turn numbers that are relevant to the current request" + description="0-indexed turn numbers that are relevant to the current request", ) context_summary: str = Field( - default="", - description="Brief summary of relevant context for the Butler" + default="", description="Brief summary of relevant context for the Butler" ) @@ -40,29 +41,28 @@ class StewardRecommendation(BaseModel): - Conversation context - Missing capabilities (if any) """ + recommended_capabilities: list[str] = Field( description="List of household member names to include (e.g., ['tatlock_core'])" ) - reasoning: str = Field( - description="Explanation of why these capabilities were recommended" - ) + reasoning: str = Field(description="Explanation of why these capabilities were recommended") estimated_complexity: Literal["simple", "moderate", "complex"] = Field( description="Complexity assessment: simple (1 tool), moderate (2-3 tools), complex (multiple tools/steps)" ) conversation_context: ConversationContext = Field( description="Contextual information from conversation history" ) - missing_capabilities: Optional[str] = Field( + missing_capabilities: str | None = Field( default=None, - description="Description of capabilities that would be helpful but aren't available" + description="Description of capabilities that would be helpful but aren't available", ) memory_context: dict[str, Any] = Field( default_factory=dict, - description="Pre-fetched user context from memory (profile, preferences)" + description="Pre-fetched user context from memory (profile, preferences)", ) enriched_query: str = Field( default="", - description="User query with auto-filled context (location, timezone) when not specified" + description="User query with auto-filled context (location, timezone) when not specified", ) def format_for_butler(self) -> str: @@ -114,8 +114,9 @@ class StewardRecommendation(BaseModel): lines.append(f" • preferences: {prefs_str}") # Add delegation instructions when expert agents are recommended - delegation_agents = [c for c in self.recommended_capabilities - if c in ("biographer", "librarian")] + delegation_agents = [ + c for c in self.recommended_capabilities if c in ("biographer", "librarian") + ] if delegation_agents: lines.append("-" * 40) lines.append("DELEGATION REQUIRED:") diff --git a/src/agents/steward/service.py b/src/agents/steward/service.py index cf4aa50..658eff5 100644 --- a/src/agents/steward/service.py +++ b/src/agents/steward/service.py @@ -7,12 +7,14 @@ and error handling. Parses plain text recommendations into structured data. Includes memory pre-fetch for user context injection. """ + import re -from typing import Any, Optional +from typing import Any from src.core.household_registry import get_household_registry from src.core.logging_config import get_logger, log_operation from src.core.memory_service import memory_service + from .agent import get_steward_agent from .schemas import ConversationContext, StewardRecommendation @@ -126,8 +128,7 @@ def _extract_complexity(text: str) -> str: def _extract_conversation_context( - text: str, - conversation_history: list[dict] + text: str, conversation_history: list[dict] ) -> ConversationContext: """ Extract conversation context analysis from text. @@ -142,13 +143,15 @@ def _extract_conversation_context( text_lower = text.lower() # Check if conversation history is referenced - has_context = bool(conversation_history) and any([ - "previous" in text_lower, - "earlier" in text_lower, - "context" in text_lower, - "turn" in text_lower, - "history" in text_lower, - ]) + has_context = bool(conversation_history) and any( + [ + "previous" in text_lower, + "earlier" in text_lower, + "context" in text_lower, + "turn" in text_lower, + "history" in text_lower, + ] + ) # Extract turn numbers if mentioned (e.g., "turn 0", "turn 1") relevant_turns = [] @@ -160,21 +163,23 @@ def _extract_conversation_context( context_summary = "" if has_context: # Extract sentence(s) mentioning context - sentences = text.split('.') - context_sentences = [s for s in sentences if any( - word in s.lower() for word in ["previous", "earlier", "context", "history"] - )] + sentences = text.split(".") + context_sentences = [ + s + for s in sentences + if any(word in s.lower() for word in ["previous", "earlier", "context", "history"]) + ] if context_sentences: context_summary = context_sentences[0].strip() return ConversationContext( has_previous_context=has_context, relevant_turns=relevant_turns, - context_summary=context_summary + context_summary=context_summary, ) -def _extract_missing_capabilities(text: str) -> Optional[str]: +def _extract_missing_capabilities(text: str) -> str | None: """ Extract missing capability notes from text. @@ -187,15 +192,16 @@ def _extract_missing_capabilities(text: str) -> Optional[str]: text_lower = text.lower() # Look for indicators of missing capabilities - if any(word in text_lower for word in [ - "missing", "unavailable", "not available", "don't have", "doesn't have" - ]): + if any( + word in text_lower + for word in ["missing", "unavailable", "not available", "don't have", "doesn't have"] + ): # Find the sentence mentioning missing capabilities - sentences = text.split('.') + sentences = text.split(".") for sentence in sentences: - if any(word in sentence.lower() for word in [ - "missing", "unavailable", "not available" - ]): + if any( + word in sentence.lower() for word in ["missing", "unavailable", "not available"] + ): return sentence.strip() return None @@ -236,7 +242,7 @@ def _build_enriched_query(user_request: str, memory_context: dict[str, Any]) -> # Check if location is needed and not specified location_keywords = ["weather", "temperature", "forecast", "nearby", "local", "here"] # Use word boundary pattern to avoid false positives like "at" in "what" - location_prepositions = [r'\bin\b', r'\bat\b', r'\bnear\b', r'\baround\b', r'\bfor\b'] + location_prepositions = [r"\bin\b", r"\bat\b", r"\bnear\b", r"\baround\b", r"\bfor\b"] location_specified = any(re.search(p, request_lower) for p in location_prepositions) if any(word in request_lower for word in location_keywords): @@ -287,32 +293,65 @@ async def _prefetch_memory_context(user_request: str) -> dict[str, Any]: profile_keys = [] # Location-related queries - if any(word in request_lower for word in [ - "weather", "temperature", "forecast", "nearby", "local", - "directions", "distance", "map", "here", - # Direct location questions - "live", "where", "home", "reside", "location", "address", - ]): + if any( + word in request_lower + for word in [ + "weather", + "temperature", + "forecast", + "nearby", + "local", + "directions", + "distance", + "map", + "here", + # Direct location questions + "live", + "where", + "home", + "reside", + "location", + "address", + ] + ): profile_keys.append("location") # Time-related queries - if any(word in request_lower for word in [ - "time", "schedule", "meeting", "appointment", "reminder", - "alarm", "when", "today", "tomorrow" - ]): + if any( + word in request_lower + for word in [ + "time", + "schedule", + "meeting", + "appointment", + "reminder", + "alarm", + "when", + "today", + "tomorrow", + ] + ): profile_keys.append("timezone") # Personal queries - if any(word in request_lower for word in [ - "my name", "who am i", "about me" - ]): + if any(word in request_lower for word in ["my name", "who am i", "about me"]): profile_keys.append("name") # Always fetch preferences if they might affect response format - include_preferences = any(word in request_lower for word in [ - "temperature", "weather", "convert", "unit", "format", - "celsius", "fahrenheit", "metric", "imperial" - ]) + include_preferences = any( + word in request_lower + for word in [ + "temperature", + "weather", + "convert", + "unit", + "format", + "celsius", + "fahrenheit", + "metric", + "imperial", + ] + ) try: return await memory_service.prefetch_context( @@ -331,7 +370,7 @@ async def _prefetch_memory_context(user_request: str) -> dict[str, Any]: async def analyze_request( user_request: str, conversation_history: list[dict], - conversation_id: Optional[str] = None, + conversation_id: str | None = None, ) -> StewardRecommendation: """ Analyze user request with full conversation context. @@ -363,7 +402,7 @@ async def analyze_request( "request_preview": user_request[:100], "conversation_id": conversation_id, "history_length": len(conversation_history), - } + }, ) as log_ctx: try: # Pre-fetch user context from memory (fast, no LLM) @@ -382,8 +421,7 @@ async def analyze_request( # Get plain text analysis from Steward analysis_text = await steward.analyze( - user_request, - conversation_history=conversation_history + user_request, conversation_history=conversation_history ) # Parse plain text into structured recommendation diff --git a/src/agents/tatlock.py b/src/agents/tatlock.py index af7702e..4d04d4b 100644 --- a/src/agents/tatlock.py +++ b/src/agents/tatlock.py @@ -6,24 +6,25 @@ The agent embodies a witty, capable British butler personality. """ import secrets -from typing import AsyncGenerator, Any +from collections.abc import AsyncGenerator from dataclasses import dataclass, field +from typing import Any from pydantic_ai import Agent, RunContext from src.agents.base import AgentInterface, OutputItem from src.agents.tatlock_core.tools import ( calculate, - get_current_datetime, calculate_time_offset, + get_current_datetime, time_difference, ) -from src.core.config import config from src.core.logging_config import get_logger from src.core.tracing import ( - start_span, end_span, get_current_span, + SpanType, add_tool_spans_from_messages, - SpanType, SpanStatus, + end_span, + start_span, ) logger = get_logger(__name__) @@ -32,6 +33,7 @@ logger = get_logger(__name__) @dataclass class ToolCallTracker: """Tracks tool calls for reporting to reasoning output.""" + calls: list[str] = field(default_factory=list) def log_call(self, message: str): @@ -239,7 +241,9 @@ class TatlockAgent(AgentInterface): # Time difference calculator @self._agent.tool - def calculate_time_difference(ctx: RunContext[ToolCallTracker], date1_str: str, date2_str: str = "now") -> str: + def calculate_time_difference( + ctx: RunContext[ToolCallTracker], date1_str: str, date2_str: str = "now" + ) -> str: """ Calculate the difference between two dates. @@ -251,7 +255,9 @@ class TatlockAgent(AgentInterface): Human-readable description of the time difference """ if ctx.deps: - ctx.deps.log_call(f"🕐 Calculating time difference between {date1_str} and {date2_str}") + ctx.deps.log_call( + f"🕐 Calculating time difference between {date1_str} and {date2_str}" + ) return time_difference(date1_str, date2_str) # NOTE: Web search has been moved to The Librarian agent. @@ -271,7 +277,7 @@ class TatlockAgent(AgentInterface): temperature: float = 1.0, max_tokens: int | None = None, stop: list[str] | None = None, - **kwargs: Any + **kwargs: Any, ) -> AsyncGenerator[OutputItem, None]: """ Generate response using PydanticAI with Ollama. @@ -305,18 +311,20 @@ class TatlockAgent(AgentInterface): type="message", id=f"msg_{generate_id()}", role="assistant", - content=[{ - "type": "output_text", - "text": "I'm afraid I didn't receive a message, sir. How may I assist you?", - "annotations": [] - }], - status="completed" + content=[ + { + "type": "output_text", + "text": "I'm afraid I didn't receive a message, sir. How may I assist you?", + "annotations": [], + } + ], + status="completed", ) return # Build message history (all messages except the last user message) # PydanticAI expects history as list of ModelRequest/ModelResponse objects - from pydantic_ai.messages import ModelRequest, ModelResponse, UserPromptPart, TextPart + from pydantic_ai.messages import ModelRequest, ModelResponse, TextPart, UserPromptPart message_history = [] for i, msg in enumerate(messages[:-1]): # All messages except the last one @@ -330,7 +338,9 @@ class TatlockAgent(AgentInterface): # Debug: Check for problematic content if '"' in content or "'" in content: - logger.debug(f"Message {i} ({role}) contains quotes. Content preview: {content[:100]}...") + logger.debug( + f"Message {i} ({role}) contains quotes. Content preview: {content[:100]}..." + ) # Convert to PydanticAI message format try: @@ -339,9 +349,7 @@ class TatlockAgent(AgentInterface): ModelRequest(parts=[UserPromptPart(content=content)]) ) elif role == "assistant": - message_history.append( - ModelResponse(parts=[TextPart(content=content)]) - ) + message_history.append(ModelResponse(parts=[TextPart(content=content)])) except Exception as e: logger.error(f"Error creating message history item {i}: {e}") logger.error(f"Problematic content: {repr(content)}") @@ -352,7 +360,9 @@ class TatlockAgent(AgentInterface): if message_history: for i, hist_msg in enumerate(message_history): msg_type = type(hist_msg).__name__ - content_preview = str(hist_msg.parts[0].content)[:50] if hist_msg.parts else "no parts" + content_preview = ( + str(hist_msg.parts[0].content)[:50] if hist_msg.parts else "no parts" + ) logger.info(f" History[{i}]: {msg_type} - {content_preview}...") # Generate reasoning output if requested @@ -362,10 +372,10 @@ class TatlockAgent(AgentInterface): id=f"reasoning_{generate_id()}", summary=[ "Analyzing your request, sir...", - "Formulating response based on available knowledge..." + "Formulating response based on available knowledge...", ], thinking="", # PydanticAI doesn't expose internal reasoning yet - status="completed" + status="completed", ) # Create a tool call tracker for this request @@ -382,7 +392,7 @@ class TatlockAgent(AgentInterface): result = await self.agent.run( user_message, message_history=message_history if message_history else None, - deps=tracker + deps=tracker, ) final_text = result.output @@ -393,7 +403,7 @@ class TatlockAgent(AgentInterface): id=f"reasoning_tools_{generate_id()}", summary=tracker.calls, thinking="", - status="completed" + status="completed", ) # Yield the complete message @@ -402,12 +412,8 @@ class TatlockAgent(AgentInterface): type="message", id=msg_id, role="assistant", - content=[{ - "type": "output_text", - "text": final_text, - "annotations": [] - }], - status="completed" + content=[{"type": "output_text", "text": final_text, "annotations": []}], + status="completed", ) except Exception as e: @@ -416,12 +422,14 @@ class TatlockAgent(AgentInterface): type="message", id=f"msg_{generate_id()}", role="assistant", - content=[{ - "type": "output_text", - "text": f"My apologies, sir. I encountered an error: {str(e)}", - "annotations": [] - }], - status="failed" + content=[ + { + "type": "output_text", + "text": f"My apologies, sir. I encountered an error: {str(e)}", + "annotations": [], + } + ], + status="failed", ) async def supports_tools(self) -> bool: @@ -490,7 +498,7 @@ class TatlockAgent(AgentInterface): enriched_message = f"{steward_note}\n\n{user_message}" # Convert message history to PydanticAI format - from pydantic_ai.messages import ModelRequest, ModelResponse, UserPromptPart, TextPart + from pydantic_ai.messages import ModelRequest, ModelResponse, TextPart, UserPromptPart pydantic_history = [] for msg in message_history: @@ -501,17 +509,14 @@ class TatlockAgent(AgentInterface): continue if role == "user": - pydantic_history.append( - ModelRequest(parts=[UserPromptPart(content=content)]) - ) + pydantic_history.append(ModelRequest(parts=[UserPromptPart(content=content)])) elif role == "assistant": - pydantic_history.append( - ModelResponse(parts=[TextPart(content=content)]) - ) + pydantic_history.append(ModelResponse(parts=[TextPart(content=content)])) # Run with scoped tools and tracker # Force tool_choice to make LLM actually call tools from src.anthropic.model_selector import get_tool_choice_settings + result = await scoped_agent.run( enriched_message, message_history=pydantic_history if pydantic_history else None, @@ -575,7 +580,7 @@ class TatlockAgent(AgentInterface): enriched_message = f"{steward_note}\n\n{user_message}" # Convert message history to PydanticAI format - from pydantic_ai.messages import ModelRequest, ModelResponse, UserPromptPart, TextPart + from pydantic_ai.messages import ModelRequest, ModelResponse, TextPart, UserPromptPart pydantic_history = [] for msg in message_history: @@ -586,13 +591,9 @@ class TatlockAgent(AgentInterface): continue if role == "user": - pydantic_history.append( - ModelRequest(parts=[UserPromptPart(content=content)]) - ) + pydantic_history.append(ModelRequest(parts=[UserPromptPart(content=content)])) elif role == "assistant": - pydantic_history.append( - ModelResponse(parts=[TextPart(content=content)]) - ) + pydantic_history.append(ModelResponse(parts=[TextPart(content=content)])) # Use run() instead of run_stream() to avoid Ollama 400 bug # with streaming + tool calls (PydanticAI issues #1292, #2256) @@ -600,7 +601,7 @@ class TatlockAgent(AgentInterface): result = await scoped_agent.run( enriched_message, message_history=pydantic_history if pydantic_history else None, - deps=tool_tracker + deps=tool_tracker, ) # Stream the final response in chunks to maintain UX @@ -608,7 +609,7 @@ class TatlockAgent(AgentInterface): chunk_size = 50 # characters per chunk for i in range(0, len(response_text), chunk_size): - yield response_text[i:i + chunk_size] + yield response_text[i : i + chunk_size] logger.info("tatlock_scoped_run_complete") @@ -643,11 +644,12 @@ class TatlockAgent(AgentInterface): from pydantic_ai.messages import ( ModelRequest, ModelResponse, - UserPromptPart, TextPart, ToolCallPart, ToolReturnPart, + UserPromptPart, ) + from src.anthropic.model_selector import get_model logger.info( @@ -663,7 +665,7 @@ class TatlockAgent(AgentInterface): SpanType.TATLOCK, metadata={ "scoped_tool_count": len(scoped_tools), - "tool_names": [getattr(t, '__name__', str(t)) for t in scoped_tools[:5]], + "tool_names": [getattr(t, "__name__", str(t)) for t in scoped_tools[:5]], }, ) @@ -690,16 +692,13 @@ class TatlockAgent(AgentInterface): continue if role == "user": - pydantic_history.append( - ModelRequest(parts=[UserPromptPart(content=content)]) - ) + pydantic_history.append(ModelRequest(parts=[UserPromptPart(content=content)])) elif role == "assistant": - pydantic_history.append( - ModelResponse(parts=[TextPart(content=content)]) - ) + pydantic_history.append(ModelResponse(parts=[TextPart(content=content)])) # Run with scoped tools and tracker from src.anthropic.model_selector import get_tool_choice_settings + result = await scoped_agent.run( enriched_message, message_history=pydantic_history if pydantic_history else None, @@ -782,7 +781,8 @@ class TatlockAgent(AgentInterface): Returns: str: Butler-toned response synthesized from all results """ - from pydantic_ai.messages import ModelRequest, ModelResponse, UserPromptPart, TextPart + from pydantic_ai.messages import ModelRequest, ModelResponse, TextPart, UserPromptPart + from src.anthropic.model_selector import get_model logger.info( @@ -849,13 +849,9 @@ class TatlockAgent(AgentInterface): continue if role == "user": - pydantic_history.append( - ModelRequest(parts=[UserPromptPart(content=content)]) - ) + pydantic_history.append(ModelRequest(parts=[UserPromptPart(content=content)])) elif role == "assistant": - pydantic_history.append( - ModelResponse(parts=[TextPart(content=content)]) - ) + pydantic_history.append(ModelResponse(parts=[TextPart(content=content)])) # Run synthesis result = await synthesis_agent.run( @@ -885,9 +881,9 @@ class TatlockAgent(AgentInterface): async def get_capabilities(self) -> dict: """Return current capabilities.""" return { - "streaming": True, # Streaming implemented - "reasoning": True, # Basic reasoning summaries - "tools": True, # Permanent tools: calculator, date/time, search - "vision": False, # Future - "audio": False, # Future + "streaming": True, # Streaming implemented + "reasoning": True, # Basic reasoning summaries + "tools": True, # Permanent tools: calculator, date/time, search + "vision": False, # Future + "audio": False, # Future } diff --git a/src/agents/tatlock_core/__init__.py b/src/agents/tatlock_core/__init__.py index ce6677a..4eb2853 100644 --- a/src/agents/tatlock_core/__init__.py +++ b/src/agents/tatlock_core/__init__.py @@ -5,14 +5,15 @@ Provides calculator and date/time capabilities. Web search has been moved to The Librarian agent. Organized as a household member with toolset and capability registration. """ + from .capability import TATLOCK_CORE_CAPABILITY, get_capability -from .toolset import get_core_tools, tatlock_core_tools from .tools import ( calculate, calculate_time_offset, get_current_datetime, time_difference, ) +from .toolset import get_core_tools, tatlock_core_tools __all__ = [ # Tools diff --git a/src/agents/tatlock_core/capability.py b/src/agents/tatlock_core/capability.py index ade3b85..9578271 100644 --- a/src/agents/tatlock_core/capability.py +++ b/src/agents/tatlock_core/capability.py @@ -4,8 +4,8 @@ Household capability definition for Tatlock's core tools. Provides the executive summary that the Steward and Butler see for coordinating household capabilities. """ -from src.core.household_registry import HouseholdCapability +from src.core.household_registry import HouseholdCapability TATLOCK_CORE_CAPABILITY = HouseholdCapability( name="tatlock_core", diff --git a/src/agents/tatlock_core/tools.py b/src/agents/tatlock_core/tools.py index 12e9ffc..474cc5b 100644 --- a/src/agents/tatlock_core/tools.py +++ b/src/agents/tatlock_core/tools.py @@ -6,13 +6,11 @@ These tools are always available to the butler agent: - Date/Time toolkit: For current time and time calculations - SearXNG search: For searching the web for current information """ + import math import re from datetime import datetime, timedelta -import httpx - -from src.core.config import config from src.core.logging_config import get_logger logger = get_logger(__name__) @@ -22,6 +20,7 @@ logger = get_logger(__name__) # Calculator Tool # ============================================================================ + def calculate(expression: str) -> str: """ Safely evaluate mathematical expressions. @@ -50,33 +49,29 @@ def calculate(expression: str) -> str: # Create safe namespace with math functions safe_dict = { # Basic math functions - 'sqrt': math.sqrt, - 'pow': math.pow, - 'abs': abs, - 'round': round, - + "sqrt": math.sqrt, + "pow": math.pow, + "abs": abs, + "round": round, # Trigonometric - 'sin': math.sin, - 'cos': math.cos, - 'tan': math.tan, - 'asin': math.asin, - 'acos': math.acos, - 'atan': math.atan, - + "sin": math.sin, + "cos": math.cos, + "tan": math.tan, + "asin": math.asin, + "acos": math.acos, + "atan": math.atan, # Logarithmic - 'log': math.log, - 'log10': math.log10, - 'log2': math.log2, - 'exp': math.exp, - + "log": math.log, + "log10": math.log10, + "log2": math.log2, + "exp": math.exp, # Other - 'ceil': math.ceil, - 'floor': math.floor, - 'factorial': math.factorial, - + "ceil": math.ceil, + "floor": math.floor, + "factorial": math.factorial, # Constants - 'pi': math.pi, - 'e': math.e, + "pi": math.pi, + "e": math.e, } # Evaluate the expression safely @@ -101,6 +96,7 @@ def calculate(expression: str) -> str: # Date/Time Toolkit # ============================================================================ + def get_current_datetime(format_str: str = "full") -> str: """ Get the current date and time. @@ -161,7 +157,7 @@ def calculate_time_offset(offset_description: str) -> str: # Parse the offset description # Pattern: "N unit(s) ago/from now" - pattern = r'(\d+)\s+(second|minute|hour|day|week|month|year)s?\s+(ago|from\s+now)' + pattern = r"(\d+)\s+(second|minute|hour|day|week|month|year)s?\s+(ago|from\s+now)" match = re.match(pattern, offset_description.lower().strip()) if not match: diff --git a/src/agents/tatlock_core/toolset.py b/src/agents/tatlock_core/toolset.py index 737f124..efe8f92 100644 --- a/src/agents/tatlock_core/toolset.py +++ b/src/agents/tatlock_core/toolset.py @@ -4,11 +4,11 @@ PydanticAI toolset for Tatlock's core tools. Converts the core tool functions into PydanticAI tool definitions that can be registered with agents and the household registry. """ + from pydantic_ai.tools import Tool from . import tools - # Create tool definitions for PydanticAI calculator_tool = Tool( function=tools.calculate, diff --git a/src/agents/tools.py b/src/agents/tools.py index 58efe3d..8e676c1 100644 --- a/src/agents/tools.py +++ b/src/agents/tools.py @@ -13,11 +13,11 @@ import math import re from datetime import datetime, timedelta - # ============================================================================ # Calculator Tool # ============================================================================ + def calculate(expression: str) -> str: """ Safely evaluate mathematical expressions. @@ -46,33 +46,29 @@ def calculate(expression: str) -> str: # Create safe namespace with math functions safe_dict = { # Basic math functions - 'sqrt': math.sqrt, - 'pow': math.pow, - 'abs': abs, - 'round': round, - + "sqrt": math.sqrt, + "pow": math.pow, + "abs": abs, + "round": round, # Trigonometric - 'sin': math.sin, - 'cos': math.cos, - 'tan': math.tan, - 'asin': math.asin, - 'acos': math.acos, - 'atan': math.atan, - + "sin": math.sin, + "cos": math.cos, + "tan": math.tan, + "asin": math.asin, + "acos": math.acos, + "atan": math.atan, # Logarithmic - 'log': math.log, - 'log10': math.log10, - 'log2': math.log2, - 'exp': math.exp, - + "log": math.log, + "log10": math.log10, + "log2": math.log2, + "exp": math.exp, # Other - 'ceil': math.ceil, - 'floor': math.floor, - 'factorial': math.factorial, - + "ceil": math.ceil, + "floor": math.floor, + "factorial": math.factorial, # Constants - 'pi': math.pi, - 'e': math.e, + "pi": math.pi, + "e": math.e, } # Evaluate the expression safely @@ -97,6 +93,7 @@ def calculate(expression: str) -> str: # Date/Time Toolkit # ============================================================================ + def get_current_datetime(format_str: str = "full") -> str: """ Get the current date and time. @@ -157,7 +154,7 @@ def calculate_time_offset(offset_description: str) -> str: # Parse the offset description # Pattern: "N unit(s) ago/from now" - pattern = r'(\d+)\s+(second|minute|hour|day|week|month|year)s?\s+(ago|from\s+now)' + pattern = r"(\d+)\s+(second|minute|hour|day|week|month|year)s?\s+(ago|from\s+now)" match = re.match(pattern, offset_description.lower().strip()) if not match: diff --git a/src/chat/router.py b/src/chat/router.py index b288c9d..53c4216 100644 --- a/src/chat/router.py +++ b/src/chat/router.py @@ -2,9 +2,10 @@ Chat completion router. OpenAI-compatible /v1/chat/completions endpoint. """ + import json import logging -from typing import AsyncGenerator +from collections.abc import AsyncGenerator from fastapi import APIRouter from starlette.responses import StreamingResponse diff --git a/src/chat/schemas.py b/src/chat/schemas.py index d781689..36ddf30 100644 --- a/src/chat/schemas.py +++ b/src/chat/schemas.py @@ -2,6 +2,7 @@ OpenAI-compatible chat completion schemas. Following OpenAI API specification for compatibility. """ + from typing import Literal from pydantic import Field @@ -11,6 +12,7 @@ from src.core.models import CustomBaseModel class ChatMessage(CustomBaseModel): """OpenAI-compatible chat message.""" + role: Literal["system", "user", "assistant"] content: str name: str | None = None @@ -18,6 +20,7 @@ class ChatMessage(CustomBaseModel): class ChatCompletionRequest(CustomBaseModel): """OpenAI-compatible chat completion request.""" + model: str = Field(..., description="Model to use for completion") messages: list[ChatMessage] = Field(..., description="List of messages") temperature: float | None = Field(default=0.7, ge=0.0, le=2.0) @@ -29,6 +32,7 @@ class ChatCompletionRequest(CustomBaseModel): class ChatCompletionChoice(CustomBaseModel): """Choice in chat completion response.""" + index: int message: ChatMessage finish_reason: str | None @@ -36,6 +40,7 @@ class ChatCompletionChoice(CustomBaseModel): class ChatCompletionUsage(CustomBaseModel): """Token usage information.""" + prompt_tokens: int completion_tokens: int total_tokens: int @@ -43,6 +48,7 @@ class ChatCompletionUsage(CustomBaseModel): class ChatCompletionResponse(CustomBaseModel): """OpenAI-compatible chat completion response.""" + id: str object: str = "chat.completion" created: int @@ -53,6 +59,7 @@ class ChatCompletionResponse(CustomBaseModel): class ChatCompletionChunkDelta(CustomBaseModel): """Delta in streaming chunk.""" + role: str | None = None content: str | None = None reasoning_content: str | None = None # For thinking/reasoning (DeepSeek R1 format) @@ -60,6 +67,7 @@ class ChatCompletionChunkDelta(CustomBaseModel): class ChatCompletionChunkChoice(CustomBaseModel): """Choice in streaming chunk.""" + index: int delta: ChatCompletionChunkDelta finish_reason: str | None = None @@ -67,6 +75,7 @@ class ChatCompletionChunkChoice(CustomBaseModel): class ChatCompletionChunk(CustomBaseModel): """OpenAI-compatible streaming chunk.""" + id: str object: str = "chat.completion.chunk" created: int diff --git a/src/chat/service.py b/src/chat/service.py index ca68796..bfec135 100644 --- a/src/chat/service.py +++ b/src/chat/service.py @@ -4,17 +4,17 @@ Chat completion service. Wrapper around Responses API that converts to Chat Completions format. Embeds reasoning in tags for Open WebUI compatibility. """ -import asyncio + import time import uuid -from typing import AsyncGenerator +from collections.abc import AsyncGenerator from src.chat import constants from src.chat.schemas import ( + ChatCompletionChoice, ChatCompletionChunk, ChatCompletionChunkChoice, ChatCompletionChunkDelta, - ChatCompletionChoice, ChatCompletionRequest, ChatCompletionResponse, ChatCompletionUsage, @@ -43,10 +43,7 @@ async def create_chat_completion( created_at = int(time.time()) # Convert Chat request to Responses request - input_messages = [ - {"role": msg.role, "content": msg.content} - for msg in request.messages - ] + input_messages = [{"role": msg.role, "content": msg.content} for msg in request.messages] response_request = ResponseRequest( model=request.model, @@ -54,7 +51,9 @@ async def create_chat_completion( reasoning={"effort": "medium", "summary": "auto"}, # Enable reasoning temperature=request.temperature or 1.0, max_output_tokens=request.max_tokens, - stop=request.stop if isinstance(request.stop, list) else ([request.stop] if request.stop else None), + stop=request.stop + if isinstance(request.stop, list) + else ([request.stop] if request.stop else None), ) # Call Responses API (will use Steward for Tatlock) @@ -118,16 +117,13 @@ async def create_chat_completion_stream( Yields: Chat completion chunks with reasoning as tags """ - from src.responses.streaming import StreamingCoordinator, StreamEventType + from src.responses.streaming import StreamEventType, StreamingCoordinator completion_id = f"chatcmpl-{uuid.uuid4().hex[:24]}" created_at = int(time.time()) # Convert Chat request to Responses request - input_messages = [ - {"role": msg.role, "content": msg.content} - for msg in request.messages - ] + input_messages = [{"role": msg.role, "content": msg.content} for msg in request.messages] response_request = ResponseRequest( model=request.model, @@ -135,7 +131,9 @@ async def create_chat_completion_stream( reasoning={"effort": "medium", "summary": "auto"}, temperature=request.temperature or 1.0, max_output_tokens=request.max_tokens, - stop=request.stop if isinstance(request.stop, list) else ([request.stop] if request.stop else None), + stop=request.stop + if isinstance(request.stop, list) + else ([request.stop] if request.stop else None), stream=True, ) diff --git a/src/core/config.py b/src/core/config.py index cfd0a40..8b36bed 100644 --- a/src/core/config.py +++ b/src/core/config.py @@ -2,6 +2,7 @@ Global application configuration. Following best practice of splitting config across domains. """ + from enum import Enum from functools import lru_cache from pathlib import Path @@ -42,6 +43,7 @@ def _get_version_from_pyproject() -> str: class Environment(str, Enum): """Application environment.""" + DEVELOPMENT = "development" PRODUCTION = "production" TESTING = "testing" @@ -54,6 +56,7 @@ class Config(BaseSettings): Loads from environment variables and .env file. Domain-specific configs should be in their respective modules. """ + model_config = SettingsConfigDict( env_file=".env", env_file_encoding="utf-8", @@ -74,143 +77,90 @@ class Config(BaseSettings): # Anthropic Configuration (Claude - cloud fallback) ANTHROPIC_API_KEY: str | None = Field( - default=None, - description="Anthropic API key for the Claude fallback backend" + default=None, description="Anthropic API key for the Claude fallback backend" ) ANTHROPIC_MODEL: str = Field( - default="claude-sonnet-5", - description="Claude model for the fallback backend" + default="claude-sonnet-5", description="Claude model for the fallback backend" ) PREFER_CLOUD_BACKEND: bool = Field( - default=False, - description="Prefer Claude over Ollama (default: local-first)" + default=False, description="Prefer Claude over Ollama (default: local-first)" ) # Ollama Configuration (local - primary backend) - OLLAMA_HOST: HttpUrl = Field( - default="http://localhost:11434", - description="Ollama server URL" - ) - OLLAMA_DEFAULT_MODEL: str = Field( - default="gemma4:e2b", - description="Default Ollama model" - ) - OLLAMA_TIMEOUT: int = Field( - default=120, - description="Ollama request timeout in seconds" - ) + OLLAMA_HOST: HttpUrl = Field(default="http://localhost:11434", description="Ollama server URL") + OLLAMA_DEFAULT_MODEL: str = Field(default="gemma4:e2b", description="Default Ollama model") + OLLAMA_TIMEOUT: int = Field(default=120, description="Ollama request timeout in seconds") STEWARD_TIMEOUT: int = Field( - default=60, - description="Steward analysis timeout in seconds (gemma4 needs ~35s warm)" + default=60, description="Steward analysis timeout in seconds (gemma4 needs ~35s warm)" ) STREAM_TIMEOUT: int = Field( - default=20, - description="Timeout for each streaming turn in seconds" + default=20, description="Timeout for each streaming turn in seconds" ) # SearXNG Configuration SEARXNG_HOST: HttpUrl = Field( default="http://searxng:8080", - description="SearXNG server URL (container name; internal port 8080)" - ) - SEARXNG_TIMEOUT: int = Field( - default=30, - description="SearXNG request timeout in seconds" + description="SearXNG server URL (container name; internal port 8080)", ) + SEARXNG_TIMEOUT: int = Field(default=30, description="SearXNG request timeout in seconds") # Redis Configuration - REDIS_HOST: str = Field( - default="localhost", - description="Redis server host" - ) - REDIS_PORT: int = Field( - default=6379, - description="Redis server port" - ) - REDIS_TIMEOUT: int = Field( - default=5, - description="Redis connection timeout in seconds" - ) + REDIS_HOST: str = Field(default="localhost", description="Redis server host") + REDIS_PORT: int = Field(default=6379, description="Redis server port") + REDIS_TIMEOUT: int = Field(default=5, description="Redis connection timeout in seconds") # Library-Desk Configuration (The Librarian backend) LIBRARIAN_TIMEOUT: int = Field( - default=180, - description="Total time budget for a librarian delegation in seconds" + default=180, description="Total time budget for a librarian delegation in seconds" ) LIBRARY_DESK_HOST: HttpUrl = Field( default="http://library-desk:8089", - description="Library-Desk API URL (container name; internal port 8089)" + description="Library-Desk API URL (container name; internal port 8089)", ) LIBRARY_DESK_API_KEY: str = Field( - default="", - description="API key for Library-Desk authentication" + default="", description="API key for Library-Desk authentication" ) LIBRARY_DESK_TIMEOUT: int = Field( - default=60, - description="Library-Desk request timeout in seconds" + default=60, description="Library-Desk request timeout in seconds" ) # Core-API Configuration (The Housekeeper backend) CORE_API_HOST: HttpUrl = Field( default="http://core-api:8083", - description="Core-API URL for Home Assistant integration (container name; internal port 8083)" - ) - CORE_API_KEY: str = Field( - default="", - description="API key for Core-API authentication" - ) - CORE_API_TIMEOUT: int = Field( - default=30, - description="Core-API request timeout in seconds" + description="Core-API URL for Home Assistant integration (container name; internal port 8083)", ) + CORE_API_KEY: str = Field(default="", description="API key for Core-API authentication") + CORE_API_TIMEOUT: int = Field(default=30, description="Core-API request timeout in seconds") # Qdrant Configuration (Memory vector storage) - QDRANT_HOST: str = Field( - default="localhost", - description="Qdrant server host" - ) - QDRANT_PORT: int = Field( - default=6333, - description="Qdrant server port" - ) + QDRANT_HOST: str = Field(default="localhost", description="Qdrant server host") + QDRANT_PORT: int = Field(default=6333, description="Qdrant server port") QDRANT_EMBEDDING_DIM: int = Field( - default=768, - description="Embedding dimension (768 for nomic-embed-text)" + default=768, description="Embedding dimension (768 for nomic-embed-text)" ) # Ollama Embedding Configuration OLLAMA_EMBEDDING_MODEL: str = Field( - default="nomic-embed-text", - description="Ollama model for embeddings" + default="nomic-embed-text", description="Ollama model for embeddings" ) # Redis Memory Database - REDIS_MEMORY_DB: int = Field( - default=1, - description="Redis database number for memory cache" - ) - REDIS_MEMORY_TTL_HOURS: int = Field( - default=24, - description="TTL for session context in hours" - ) + REDIS_MEMORY_DB: int = Field(default=1, description="Redis database number for memory cache") + REDIS_MEMORY_TTL_HOURS: int = Field(default=24, description="TTL for session context in hours") # Logging LOG_LEVEL: str | None = Field( - default=None, - description="Logging level (auto-set based on environment if not specified)" + default=None, description="Logging level (auto-set based on environment if not specified)" ) # User Configuration DEFAULT_USER: str | None = Field( default=None, - description="Default user for single-user setup (auto-set based on environment if not specified)" + description="Default user for single-user setup (auto-set based on environment if not specified)", ) # CORS - CORS_ORIGINS: list[str] = Field( - default=["*"], - description="Allowed CORS origins" - ) + CORS_ORIGINS: list[str] = Field(default=["*"], description="Allowed CORS origins") CORS_ALLOW_CREDENTIALS: bool = True CORS_ALLOW_METHODS: list[str] = ["*"] CORS_ALLOW_HEADERS: list[str] = ["*"] @@ -235,8 +185,7 @@ class Config(BaseSettings): if ( self.ENVIRONMENT != Environment.PRODUCTION and self.DEFAULT_USER is not None - and sanitize_user_id(self.DEFAULT_USER) - == sanitize_user_id(PRODUCTION_TENANT) + and sanitize_user_id(self.DEFAULT_USER) == sanitize_user_id(PRODUCTION_TENANT) ): raise ValueError( f"Refusing to start: ENVIRONMENT={self.ENVIRONMENT.value} is " @@ -299,8 +248,7 @@ class Config(BaseSettings): return self.DEFAULT_USER or PRODUCTION_TENANT if self.DEFAULT_USER is not None and ( - self.DEFAULT_USER == TEST_TENANT - or self.DEFAULT_USER.startswith(TEST_TENANT_PREFIX) + self.DEFAULT_USER == TEST_TENANT or self.DEFAULT_USER.startswith(TEST_TENANT_PREFIX) ): return self.DEFAULT_USER return TEST_TENANT diff --git a/src/core/context.py b/src/core/context.py index 811dbe1..eebf01f 100644 --- a/src/core/context.py +++ b/src/core/context.py @@ -16,6 +16,7 @@ Usage: from src.core.context import get_user user = get_user() # Returns current request's user """ + from contextvars import ContextVar @@ -28,6 +29,7 @@ def get_default_user() -> str: """ # Import here to avoid circular dependency from src.core.config import config + return config.effective_default_user @@ -36,9 +38,7 @@ def get_default_user() -> str: # and resolve the real default in get_user() _USER_NOT_SET = "__user_not_set__" current_user: ContextVar[str] = ContextVar("current_user", default=_USER_NOT_SET) -current_conversation: ContextVar[str | None] = ContextVar( - "current_conversation", default=None -) +current_conversation: ContextVar[str | None] = ContextVar("current_conversation", default=None) def apply_tenant_guard(user: str) -> str: @@ -60,9 +60,8 @@ def apply_tenant_guard(user: str) -> str: from src.core.config import PRODUCTION_TENANT, TEST_TENANT, Environment, config from src.core.multi_tenancy import sanitize_user_id - if ( - config.ENVIRONMENT != Environment.PRODUCTION - and sanitize_user_id(user) == sanitize_user_id(PRODUCTION_TENANT) + if config.ENVIRONMENT != Environment.PRODUCTION and sanitize_user_id(user) == sanitize_user_id( + PRODUCTION_TENANT ): from src.core.logging_config import get_logger diff --git a/src/core/embeddings.py b/src/core/embeddings.py index 4957ea4..8f7c705 100644 --- a/src/core/embeddings.py +++ b/src/core/embeddings.py @@ -8,7 +8,6 @@ Provides async embedding operations via Ollama API: Adapted from library-desk patterns. """ -from typing import Optional import httpx diff --git a/src/core/exceptions.py b/src/core/exceptions.py index 66e99d0..c38d31c 100644 --- a/src/core/exceptions.py +++ b/src/core/exceptions.py @@ -2,12 +2,13 @@ Global exception definitions. Domain-specific exceptions should be in their respective modules. """ + from typing import Any class AppException(Exception): """Base exception for all application errors.""" - + def __init__( self, message: str = "An error occurred", @@ -22,26 +23,26 @@ class AppException(Exception): class OllamaConnectionError(AppException): """Raised when cannot connect to Ollama service.""" - + def __init__(self, message: str = "Cannot connect to Ollama service"): super().__init__(message=message, status_code=503) class OllamaTimeoutError(AppException): """Raised when Ollama request times out.""" - + def __init__(self, message: str = "Ollama request timed out"): super().__init__(message=message, status_code=504) class ModelNotFoundError(AppException): """Raised when requested model is not available.""" - + def __init__(self, model_name: str): super().__init__( message=f"Model '{model_name}' not found", status_code=404, - details={"model": model_name} + details={"model": model_name}, ) diff --git a/src/core/household_registry.py b/src/core/household_registry.py index 780dc52..a36bc22 100644 --- a/src/core/household_registry.py +++ b/src/core/household_registry.py @@ -5,10 +5,10 @@ Provides centralized registry of household members (agents) with their capabilities and tools. Supports two-tier abstraction: executive summaries for coordination and full toolsets for execution. """ -from typing import Any, Optional + +from typing import Any from pydantic import BaseModel, ConfigDict -from pydantic_ai import Agent from .logging_config import get_logger @@ -22,6 +22,7 @@ class HouseholdCapability(BaseModel): This is what the Steward and Butler see for coordination. High-level description without implementation details. """ + name: str # Unique identifier: "tatlock_core", "librarian", "developer" role: str # Display name: "Butler's Core Tools", "The Librarian" category: str # "core", "research", "technical", "automation" @@ -38,11 +39,12 @@ class HouseholdMember(BaseModel): Contains both the executive summary (for coordination) and implementation details (tools/agent). """ + model_config = ConfigDict(arbitrary_types_allowed=True) capability: HouseholdCapability tools: list[Any] # PydanticAI tool definitions (any type since Tool is a dataclass) - agent: Optional[Any] = None # For expert agents (Phase 4) + agent: Any | None = None # For expert agents (Phase 4) class HouseholdRegistry: @@ -65,7 +67,7 @@ class HouseholdRegistry: name: str, capability: HouseholdCapability, tools: list[Any], - agent: Optional[Any] = None, + agent: Any | None = None, ) -> None: """ Register a household member. @@ -95,9 +97,7 @@ class HouseholdRegistry: ... ) """ if name != capability.name: - raise ValueError( - f"Name mismatch: '{name}' != '{capability.name}'" - ) + raise ValueError(f"Name mismatch: '{name}' != '{capability.name}'") self._members[name] = HouseholdMember( capability=capability, @@ -132,7 +132,7 @@ class HouseholdRegistry: role=member.capability.role, ) - def get_member(self, name: str) -> Optional[HouseholdMember]: + def get_member(self, name: str) -> HouseholdMember | None: """ Get full household member specification. diff --git a/src/core/logging_config.py b/src/core/logging_config.py index f8a26df..f00900d 100644 --- a/src/core/logging_config.py +++ b/src/core/logging_config.py @@ -4,12 +4,14 @@ Structured logging configuration using structlog. Deeply integrates with FastAPI/uvicorn's built-in logging to provide seamless structured logs across the entire application stack. """ + import logging import logging.config import sys +from collections.abc import AsyncIterator from contextlib import asynccontextmanager -from datetime import datetime, timezone -from typing import Any, AsyncIterator +from datetime import UTC, datetime +from typing import Any import structlog from structlog.types import EventDict, Processor @@ -19,7 +21,7 @@ from .config import config def add_timestamp(logger: Any, method_name: str, event_dict: EventDict) -> EventDict: """Add ISO 8601 timestamp to log entries.""" - event_dict["timestamp"] = datetime.now(timezone.utc).isoformat() + event_dict["timestamp"] = datetime.now(UTC).isoformat() return event_dict @@ -41,11 +43,28 @@ def extract_from_record(logger: Any, method_name: str, event_dict: EventDict) -> # Extract custom fields from record for key, value in record.__dict__.items(): if key not in { - "name", "msg", "args", "created", "filename", "funcName", - "levelname", "levelno", "lineno", "module", "msecs", - "message", "pathname", "process", "processName", "relativeCreated", - "thread", "threadName", "exc_info", "exc_text", "stack_info", - "taskName" + "name", + "msg", + "args", + "created", + "filename", + "funcName", + "levelname", + "levelno", + "lineno", + "module", + "msecs", + "message", + "pathname", + "process", + "processName", + "relativeCreated", + "thread", + "threadName", + "exc_info", + "exc_text", + "stack_info", + "taskName", }: event_dict[key] = value @@ -164,7 +183,7 @@ def get_logger(name: str) -> structlog.stdlib.BoundLogger: async def log_operation( operation: str, initial_context: dict[str, Any] | None = None, - logger_name: str = "tatlock.operations" + logger_name: str = "tatlock.operations", ) -> AsyncIterator[dict[str, Any]]: """ Context manager for automatic operation timing and logging. @@ -187,21 +206,21 @@ async def log_operation( context = initial_context or {} context["operation"] = operation - start_time = datetime.now(timezone.utc) + start_time = datetime.now(UTC) logger.info("operation_started", **context) try: yield context # Success case - duration = (datetime.now(timezone.utc) - start_time).total_seconds() + duration = (datetime.now(UTC) - start_time).total_seconds() context["duration_seconds"] = duration context["success"] = True logger.info("operation_completed", **context) except Exception as e: # Error case - duration = (datetime.now(timezone.utc) - start_time).total_seconds() + duration = (datetime.now(UTC) - start_time).total_seconds() context["duration_seconds"] = duration context["success"] = False context["error"] = str(e) @@ -228,7 +247,8 @@ def get_uvicorn_log_config() -> dict[str, Any]: "()": structlog.stdlib.ProcessorFormatter, "processors": [ structlog.stdlib.ProcessorFormatter.remove_processors_meta, - structlog.processors.JSONRenderer() if config.log_format == "json" + structlog.processors.JSONRenderer() + if config.log_format == "json" else structlog.dev.ConsoleRenderer(colors=True), ], }, diff --git a/src/core/memory_cache.py b/src/core/memory_cache.py index 7b0fe5a..e963f4a 100644 --- a/src/core/memory_cache.py +++ b/src/core/memory_cache.py @@ -8,6 +8,7 @@ Provides short-term memory storage with TTL: Uses Redis DB 1. """ + import json from typing import Any @@ -15,7 +16,7 @@ import redis.asyncio as redis from .config import config from .logging_config import get_logger -from .multi_tenancy import get_session_key, get_entities_key +from .multi_tenancy import get_entities_key, get_session_key logger = get_logger(__name__) diff --git a/src/core/memory_service.py b/src/core/memory_service.py index bff8e22..8b95e8f 100644 --- a/src/core/memory_service.py +++ b/src/core/memory_service.py @@ -21,14 +21,14 @@ Usage: # Get session context ctx = await memory_service.get_session_context(conversation_id) """ -from datetime import datetime, timezone + +from datetime import UTC, datetime from enum import Enum from typing import Any from pydantic import BaseModel, Field -from .config import config -from .context import get_user, get_conversation_id +from .context import get_conversation_id, get_user from .embeddings import get_embedding_client from .logging_config import get_logger from .memory_cache import get_memory_cache @@ -40,22 +40,24 @@ logger = get_logger(__name__) class MemoryType(str, Enum): """Types of memories stored in Qdrant.""" - USER_PROFILE = "user_profile" # Name, location, timezone - PREFERENCE = "preference" # Units, language, theme - LEARNED_FACT = "learned_fact" # "My car is a Tesla" + + USER_PROFILE = "user_profile" # Name, location, timezone + PREFERENCE = "preference" # Units, language, theme + LEARNED_FACT = "learned_fact" # "My car is a Tesla" class MemoryRecord(BaseModel): """A memory record stored in Qdrant.""" + id: str type: MemoryType - key: str # e.g., "location", "timezone", "car" - value: str # The actual content + key: str # e.g., "location", "timezone", "car" + value: str # The actual content keywords: list[str] = Field(default_factory=list) - importance: float = 0.5 # 0.0 - 1.0 - source: str = "explicit" # "explicit" | "inferred" | "conversation" - created_at: str = Field(default_factory=lambda: datetime.now(timezone.utc).isoformat()) - updated_at: str = Field(default_factory=lambda: datetime.now(timezone.utc).isoformat()) + importance: float = 0.5 # 0.0 - 1.0 + source: str = "explicit" # "explicit" | "inferred" | "conversation" + created_at: str = Field(default_factory=lambda: datetime.now(UTC).isoformat()) + updated_at: str = Field(default_factory=lambda: datetime.now(UTC).isoformat()) class MemoryService: @@ -528,7 +530,7 @@ class MemoryService: "keywords": keywords, "importance": importance, "source": source, - "updated_at": datetime.now(timezone.utc).isoformat(), + "updated_at": datetime.now(UTC).isoformat(), } result = await self.qdrant.upsert_memory( diff --git a/src/core/models.py b/src/core/models.py index 7329f44..6f3b739 100644 --- a/src/core/models.py +++ b/src/core/models.py @@ -2,6 +2,7 @@ Custom Pydantic base models for consistent serialization. Following best practice of having a global base model. """ + from datetime import datetime from typing import Any @@ -17,12 +18,13 @@ def datetime_to_iso_str(dt: datetime) -> str: class CustomBaseModel(BaseModel): """ Custom base model with consistent configuration. - + All domain models should inherit from this for: - Consistent JSON serialization - Timezone-aware datetime handling - Alias population support """ + model_config = ConfigDict( json_encoders={datetime: datetime_to_iso_str}, populate_by_name=True, @@ -30,14 +32,13 @@ class CustomBaseModel(BaseModel): validate_assignment=True, arbitrary_types_allowed=True, ) - + def serializable_dict(self, **kwargs: Any) -> dict[str, Any]: """ Return dict with only JSON-serializable fields. - + Useful for logging and debugging. """ return jsonable_encoder( - self.model_dump(**kwargs), - custom_encoder={datetime: datetime_to_iso_str} + self.model_dump(**kwargs), custom_encoder={datetime: datetime_to_iso_str} ) diff --git a/src/core/multi_tenancy.py b/src/core/multi_tenancy.py index cd47b4d..086943f 100644 --- a/src/core/multi_tenancy.py +++ b/src/core/multi_tenancy.py @@ -7,6 +7,7 @@ Provides utilities for user namespace management across: Adapted from library-desk patterns. """ + import re @@ -39,13 +40,13 @@ def sanitize_user_id(user_id: str) -> str: sanitized = sanitized.replace(".", "_") # Replace any non-alphanumeric characters with underscores - sanitized = re.sub(r'[^a-z0-9_]', '_', sanitized) + sanitized = re.sub(r"[^a-z0-9_]", "_", sanitized) # Remove consecutive underscores - sanitized = re.sub(r'_+', '_', sanitized) + sanitized = re.sub(r"_+", "_", sanitized) # Remove leading/trailing underscores - sanitized = sanitized.strip('_') + sanitized = sanitized.strip("_") return sanitized @@ -141,7 +142,7 @@ def validate_user_id(user_id: str) -> bool: return False # Must contain at least one alphanumeric character - if not re.search(r'[a-zA-Z0-9]', user_id): + if not re.search(r"[a-zA-Z0-9]", user_id): return False return True diff --git a/src/core/preprocessing.py b/src/core/preprocessing.py index 8a37e17..c9be736 100644 --- a/src/core/preprocessing.py +++ b/src/core/preprocessing.py @@ -3,15 +3,16 @@ Request preprocessing pipeline. Analyzes requests via the Steward and creates scoped toolsets for Tatlock. """ + from dataclasses import dataclass from datetime import datetime -from typing import Any, Optional +from typing import Any from src.agents.steward import analyze_request, format_steward_note from src.agents.steward.schemas import StewardRecommendation from src.core.household_registry import get_household_registry from src.core.logging_config import get_logger -from src.core.tracing import trace_span, SpanType +from src.core.tracing import SpanType, trace_span logger = get_logger(__name__) @@ -45,6 +46,7 @@ class EnrichedRequest: recommendation: Full Steward recommendation steward_reasoning: Plain text reasoning for streaming to user """ + original_request: str steward_note: str scoped_tools: list[Any] # PydanticAI tool definitions @@ -55,7 +57,7 @@ class EnrichedRequest: async def preprocess_request( user_request: str, conversation_history: list[dict], - conversation_id: Optional[str] = None, + conversation_id: str | None = None, ) -> EnrichedRequest: """ Analyze request via Steward and prepare scoped context for Tatlock. @@ -111,12 +113,14 @@ async def preprocess_request( # Update span with results if span: - span.metadata.update({ - "recommended_capabilities": recommendation.recommended_capabilities, - "complexity": recommendation.estimated_complexity, - "has_memory_context": bool(recommendation.memory_context), - "has_conversation_context": recommendation.conversation_context.has_previous_context, - }) + span.metadata.update( + { + "recommended_capabilities": recommendation.recommended_capabilities, + "complexity": recommendation.estimated_complexity, + "has_memory_context": bool(recommendation.memory_context), + "has_conversation_context": recommendation.conversation_context.has_previous_context, + } + ) span.details["reasoning"] = recommendation.reasoning if recommendation.enriched_query: span.details["enriched_query"] = recommendation.enriched_query @@ -128,9 +132,7 @@ async def preprocess_request( # Uses agent-as-tool pattern: expert agents get delegation wrappers, # core tools are returned directly registry = get_household_registry() - scoped_tools = registry.get_delegation_tools( - recommendation.recommended_capabilities - ) + scoped_tools = registry.get_delegation_tools(recommendation.recommended_capabilities) logger.info( "preprocessing_complete", diff --git a/src/core/qdrant.py b/src/core/qdrant.py index 162a92f..3509733 100644 --- a/src/core/qdrant.py +++ b/src/core/qdrant.py @@ -8,8 +8,9 @@ Provides async operations for storing and retrieving memory embeddings: Adapted from library-desk patterns. """ + from typing import Any -from uuid import uuid4, uuid5, NAMESPACE_DNS +from uuid import NAMESPACE_DNS, uuid4, uuid5 from qdrant_client import QdrantClient from qdrant_client.http import models as qdrant_models diff --git a/src/core/router.py b/src/core/router.py index 164380f..a774b20 100644 --- a/src/core/router.py +++ b/src/core/router.py @@ -1,6 +1,7 @@ """ Core router for health and root endpoints. """ + import logging from fastapi import APIRouter @@ -16,7 +17,7 @@ router = APIRouter(tags=["core"]) async def health_check() -> dict[str, str]: """ Health check endpoint. - + Returns: Health status """ @@ -27,7 +28,7 @@ async def health_check() -> dict[str, str]: async def root() -> dict[str, str]: """ Root endpoint. - + Returns: API information """ diff --git a/src/core/startup.py b/src/core/startup.py index b7fa8f9..ac2aa61 100644 --- a/src/core/startup.py +++ b/src/core/startup.py @@ -5,6 +5,7 @@ Handles initialization of household registry and other startup tasks. This module should be called during application startup to register all household members. """ + from src.agents.biographer import register_biographer from src.agents.housekeeper import register_housekeeper from src.agents.librarian import register_librarian diff --git a/src/core/tool_tracking.py b/src/core/tool_tracking.py index b42410d..052d726 100644 --- a/src/core/tool_tracking.py +++ b/src/core/tool_tracking.py @@ -4,7 +4,6 @@ Tool call tracking. Tracks which tools are recommended by the Steward versus which tools are actually used by Tatlock for debugging and analysis. """ -from typing import Optional from src.core.logging_config import get_logger @@ -19,11 +18,7 @@ class ToolCallTracker: to measure recommendation accuracy. """ - def __init__( - self, - recommended_capabilities: list[str], - conversation_id: Optional[str] = None - ): + def __init__(self, recommended_capabilities: list[str], conversation_id: str | None = None): """ Initialize tool call tracker. @@ -95,9 +90,7 @@ class ToolCallTracker: tools that were recommended but never used. """ # Normalize actual tool names to capabilities for comparison - used_capabilities = { - self._extract_capability(tool) for tool in self.actual_calls.keys() - } + used_capabilities = {self._extract_capability(tool) for tool in self.actual_calls.keys()} # Find tools that were recommended but not used unused_tools = self.recommended_capabilities - used_capabilities @@ -128,9 +121,7 @@ class ToolCallTracker: """ total_calls = sum(len(durations) for durations in self.actual_calls.values()) # Normalize actual tool names to capabilities for comparison - used_capabilities = { - self._extract_capability(tool) for tool in self.actual_calls.keys() - } + used_capabilities = {self._extract_capability(tool) for tool in self.actual_calls.keys()} unused = self.recommended_capabilities - used_capabilities return { @@ -139,12 +130,8 @@ class ToolCallTracker: "tools_unused": list(unused), "total_calls": total_calls, "accuracy": { - "recommended_and_used": len( - self.recommended_capabilities & used_capabilities - ), + "recommended_and_used": len(self.recommended_capabilities & used_capabilities), "recommended_but_unused": len(unused), - "not_recommended_but_used": len( - used_capabilities - self.recommended_capabilities - ), + "not_recommended_but_used": len(used_capabilities - self.recommended_capabilities), }, } diff --git a/src/core/tracing.py b/src/core/tracing.py index ce56396..0d5092c 100644 --- a/src/core/tracing.py +++ b/src/core/tracing.py @@ -9,16 +9,16 @@ Enable via DEBUG=true environment variable. Traces are written to logs/traces/{trace_id}.json View with logs/traces/viewer.html """ -from contextlib import asynccontextmanager -from contextvars import ContextVar -from dataclasses import dataclass, field -from datetime import datetime, timezone -from enum import Enum -from pathlib import Path -from typing import Any import json import secrets +from contextlib import asynccontextmanager +from contextvars import ContextVar +from dataclasses import dataclass, field +from datetime import UTC, datetime +from enum import Enum +from pathlib import Path +from typing import Any from src.core.logging_config import get_logger @@ -27,6 +27,7 @@ logger = get_logger(__name__) class SpanType(str, Enum): """Types of traced operations.""" + ROUTER = "router" STEWARD = "steward" TATLOCK = "tatlock" @@ -36,6 +37,7 @@ class SpanType(str, Enum): class SpanStatus(str, Enum): """Span completion status.""" + OK = "ok" ERROR = "error" @@ -43,6 +45,7 @@ class SpanStatus(str, Enum): @dataclass class Span: """A single traced operation.""" + span_id: str name: str type: SpanType @@ -88,6 +91,7 @@ class Span: @dataclass class Trace: """Complete trace of a request.""" + trace_id: str conversation_id: str | None user: str @@ -116,7 +120,9 @@ class Trace: "conversation_id": self.conversation_id, "user": self.user, "timestamp": self.timestamp.isoformat(), - "total_duration_ms": round(self.total_duration_ms, 2) if self.total_duration_ms else None, + "total_duration_ms": round(self.total_duration_ms, 2) + if self.total_duration_ms + else None, "status": self.status, "request": self.request, "response": self.response, @@ -132,6 +138,7 @@ _current_span: ContextVar[Span | None] = ContextVar("current_span", default=None def tracing_enabled() -> bool: """Check if tracing is enabled (requires DEBUG=true).""" from src.core.config import config + return config.DEBUG @@ -163,7 +170,7 @@ def start_trace( trace_id=_generate_id("trace_"), conversation_id=conversation_id, user=user, - timestamp=datetime.now(timezone.utc), + timestamp=datetime.now(UTC), request=request, ) _current_trace.set(trace) @@ -209,7 +216,7 @@ def start_span( span_id=_generate_id("span_"), name=name, type=span_type, - start_time=datetime.now(timezone.utc), + start_time=datetime.now(UTC), parent_id=parent.span_id if parent else None, metadata=metadata or {}, details=details or {}, @@ -254,7 +261,7 @@ def end_span( if not span: return - span.end_time = datetime.now(timezone.utc) + span.end_time = datetime.now(UTC) span.status = status if error: span.error = error @@ -405,7 +412,7 @@ def add_tool_spans_from_messages(messages: list[Any], parent_span: Span | None = if isinstance(part, ToolCallPart): tool_calls[part.tool_call_id] = { "name": part.tool_name, - "args": part.args if hasattr(part, 'args') else {}, + "args": part.args if hasattr(part, "args") else {}, } elif isinstance(msg, ModelRequest): for part in msg.parts: @@ -418,7 +425,7 @@ def add_tool_spans_from_messages(messages: list[Any], parent_span: Span | None = name=tool_info["name"], type=SpanType.TOOL, start_time=parent_span.start_time, # Approximate - end_time=parent_span.end_time or datetime.now(timezone.utc), + end_time=parent_span.end_time or datetime.now(UTC), parent_id=parent_span.span_id, status=SpanStatus.OK, metadata={ @@ -427,7 +434,9 @@ def add_tool_spans_from_messages(messages: list[Any], parent_span: Span | None = }, details={ "args": tool_info.get("args", {}), - "result": part.content[:2000] if isinstance(part.content, str) else str(part.content)[:2000], + "result": part.content[:2000] + if isinstance(part.content, str) + else str(part.content)[:2000], }, ) parent_span.children.append(span.span_id) diff --git a/src/core/tracing_router.py b/src/core/tracing_router.py index 9d3a00f..5fbc1c0 100644 --- a/src/core/tracing_router.py +++ b/src/core/tracing_router.py @@ -4,6 +4,8 @@ Trace viewer router. Serves the trace viewer UI and trace files when tracing is enabled. Only available when DEBUG=true. """ + +from datetime import UTC from pathlib import Path from fastapi import APIRouter, HTTPException @@ -66,12 +68,12 @@ async def list_traces( return {"traces": [], "total": 0} import json - from datetime import datetime, timezone, timedelta + from datetime import datetime, timedelta # Calculate cutoff time if filtering by time cutoff_time = None if since_minutes: - cutoff_time = datetime.now(timezone.utc) - timedelta(minutes=since_minutes) + cutoff_time = datetime.now(UTC) - timedelta(minutes=since_minutes) # Get all trace files, sorted by modification time (newest first) trace_files = sorted( @@ -93,7 +95,7 @@ async def list_traces( trace_timestamp = data.get("timestamp") if cutoff_time and trace_timestamp: try: - ts = datetime.fromisoformat(trace_timestamp.replace('Z', '+00:00')) + ts = datetime.fromisoformat(trace_timestamp.replace("Z", "+00:00")) if ts < cutoff_time: continue except (ValueError, TypeError): @@ -109,15 +111,17 @@ async def list_traces( if search and search.lower() not in request_preview.lower(): continue - traces.append({ - "trace_id": data.get("trace_id"), - "timestamp": trace_timestamp, - "user": data.get("user"), - "status": trace_status, - "total_duration_ms": data.get("total_duration_ms"), - "span_count": len(data.get("spans", [])), - "request_preview": request_preview[:100], - }) + traces.append( + { + "trace_id": data.get("trace_id"), + "timestamp": trace_timestamp, + "user": data.get("user"), + "status": trace_status, + "total_duration_ms": data.get("total_duration_ms"), + "span_count": len(data.get("spans", [])), + "request_preview": request_preview[:100], + } + ) except Exception as e: logger.warning("trace_list_parse_error", path=str(path), error=str(e)) @@ -145,6 +149,7 @@ async def get_trace(trace_id: str): try: import json + with open(trace_path) as f: data = json.load(f) return JSONResponse(content=data) diff --git a/src/main.py b/src/main.py index 1f4bf4f..ab79086 100644 --- a/src/main.py +++ b/src/main.py @@ -9,8 +9,9 @@ Main responsibilities: - Router registration - Lifecycle management """ + +from collections.abc import AsyncGenerator from contextlib import asynccontextmanager -from typing import AsyncGenerator from fastapi import FastAPI, Request, status from fastapi.exceptions import RequestValidationError @@ -64,7 +65,7 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: def create_application() -> FastAPI: """ Application factory. - + Creates and configures the FastAPI application. Following best practice of using factory pattern. """ @@ -75,7 +76,7 @@ def create_application() -> FastAPI: lifespan=lifespan, debug=config.DEBUG, ) - + # Add middleware application.add_middleware( CORSMiddleware, @@ -84,10 +85,10 @@ def create_application() -> FastAPI: allow_methods=config.CORS_ALLOW_METHODS, allow_headers=config.CORS_ALLOW_HEADERS, ) - + # Register exception handlers register_exception_handlers(application) - + # Include routers application.include_router(core_router) # Health and root endpoints application.include_router(chat_router, prefix=config.API_PREFIX) @@ -105,10 +106,10 @@ def create_application() -> FastAPI: def register_exception_handlers(application: FastAPI) -> None: """ Register global exception handlers. - + Provides consistent error responses compatible with OpenAI API. """ - + @application.exception_handler(AppException) async def app_exception_handler( request: Request, diff --git a/src/models/router.py b/src/models/router.py index 8ddb374..c51e579 100644 --- a/src/models/router.py +++ b/src/models/router.py @@ -2,6 +2,7 @@ Models router. OpenAI-compatible /v1/models endpoint. """ + import logging from fastapi import APIRouter @@ -18,9 +19,9 @@ router = APIRouter(prefix="/models", tags=["models"]) async def list_models() -> ModelsResponse: """ List available models (OpenAI-compatible). - + Currently returns mock model list. - + Returns: List of available models """ diff --git a/src/models/schemas.py b/src/models/schemas.py index bb450f8..4ded972 100644 --- a/src/models/schemas.py +++ b/src/models/schemas.py @@ -1,11 +1,13 @@ """ OpenAI-compatible models schemas. """ + from src.core.models import CustomBaseModel class Model(CustomBaseModel): """OpenAI-compatible model object.""" + id: str object: str = "model" created: int @@ -14,5 +16,6 @@ class Model(CustomBaseModel): class ModelsResponse(CustomBaseModel): """OpenAI-compatible models list response.""" + object: str = "list" data: list[Model] diff --git a/src/ollama/client.py b/src/ollama/client.py index c4ed984..aff27f5 100644 --- a/src/ollama/client.py +++ b/src/ollama/client.py @@ -2,8 +2,10 @@ Ollama HTTP client. Handles all communication with the Ollama service. """ + import logging -from typing import Any, AsyncGenerator +from collections.abc import AsyncGenerator +from typing import Any import httpx from httpx import ConnectError, TimeoutException @@ -22,14 +24,14 @@ logger = logging.getLogger(__name__) class OllamaClient: """ Async client for Ollama API. - + Follows best practice of using async for I/O operations. """ - + def __init__(self, base_url: str | None = None, timeout: int | None = None): """ Initialize Ollama client. - + Args: base_url: Ollama server URL (defaults to config) timeout: Request timeout in seconds (defaults to config) @@ -37,7 +39,7 @@ class OllamaClient: self.base_url = base_url or str(config.OLLAMA_HOST) self.timeout = timeout or config.OLLAMA_TIMEOUT self._client: httpx.AsyncClient | None = None - + async def __aenter__(self) -> "OllamaClient": """Async context manager entry.""" self._client = httpx.AsyncClient( @@ -45,32 +47,32 @@ class OllamaClient: timeout=self.timeout, ) return self - + async def __aexit__(self, *args: Any) -> None: """Async context manager exit.""" if self._client: await self._client.aclose() - + async def chat( self, request: OllamaChatRequest, ) -> OllamaChatResponse: """ Send chat request to Ollama (non-streaming). - + Args: request: Chat request with model and messages - + Returns: Complete chat response - + Raises: OllamaConnectionError: Cannot connect to Ollama OllamaTimeoutError: Request timed out """ if not self._client: raise RuntimeError("Client not initialized. Use async with context.") - + try: response = await self._client.post( "/api/chat", @@ -78,38 +80,38 @@ class OllamaClient: ) response.raise_for_status() return OllamaChatResponse(**response.json()) - + except ConnectError as e: logger.error(f"Cannot connect to Ollama at {self.base_url}: {e}") raise OllamaConnectionError() from e - + except TimeoutException as e: logger.error(f"Ollama request timed out after {self.timeout}s: {e}") raise OllamaTimeoutError() from e - + async def chat_stream( self, request: OllamaChatRequest, ) -> AsyncGenerator[dict[str, Any], None]: """ Send streaming chat request to Ollama. - + Args: request: Chat request with stream=True - + Yields: Streaming response chunks - + Raises: OllamaConnectionError: Cannot connect to Ollama OllamaTimeoutError: Request timed out """ if not self._client: raise RuntimeError("Client not initialized. Use async with context.") - + # Ensure streaming is enabled request.stream = True - + try: async with self._client.stream( "POST", @@ -120,48 +122,49 @@ class OllamaClient: async for line in response.aiter_lines(): if line.strip(): import json + yield json.loads(line) - + except ConnectError as e: logger.error(f"Cannot connect to Ollama at {self.base_url}: {e}") raise OllamaConnectionError() from e - + except TimeoutException as e: logger.error(f"Ollama request timed out after {self.timeout}s: {e}") raise OllamaTimeoutError() from e - + async def list_models(self) -> OllamaModelsResponse: """ List available models from Ollama. - + Returns: List of available models - + Raises: OllamaConnectionError: Cannot connect to Ollama """ if not self._client: raise RuntimeError("Client not initialized. Use async with context.") - + try: response = await self._client.get("/api/tags") response.raise_for_status() return OllamaModelsResponse(**response.json()) - + except ConnectError as e: logger.error(f"Cannot connect to Ollama at {self.base_url}: {e}") raise OllamaConnectionError() from e - + async def health_check(self) -> bool: """ Check if Ollama service is healthy. - + Returns: True if healthy, False otherwise """ if not self._client: raise RuntimeError("Client not initialized. Use async with context.") - + try: response = await self._client.get("/") return response.status_code == 200 @@ -174,7 +177,7 @@ class OllamaClient: async def get_ollama_client() -> AsyncGenerator[OllamaClient, None]: """ FastAPI dependency to provide Ollama client. - + Follows best practice of dependency injection. """ async with OllamaClient() as client: diff --git a/src/ollama/provider.py b/src/ollama/provider.py index 1ec7fc1..95c0e1d 100644 --- a/src/ollama/provider.py +++ b/src/ollama/provider.py @@ -5,6 +5,7 @@ Ollama's OpenAI-compatible API rejects messages with `content: null`, which PydanticAI sends for assistant messages that only contain tool calls. This provider sanitizes messages to use empty strings instead of null. """ + from typing import Any from openai import AsyncOpenAI diff --git a/src/ollama/schemas.py b/src/ollama/schemas.py index 78010ee..fe7e8b5 100644 --- a/src/ollama/schemas.py +++ b/src/ollama/schemas.py @@ -2,6 +2,7 @@ Ollama API schemas. Internal models for Ollama API communication. """ + from typing import Any from src.core.models import CustomBaseModel @@ -9,12 +10,14 @@ from src.core.models import CustomBaseModel class OllamaMessage(CustomBaseModel): """Message format for Ollama API.""" + role: str content: str class OllamaChatRequest(CustomBaseModel): """Chat request to Ollama API.""" + model: str messages: list[OllamaMessage] stream: bool = False @@ -23,6 +26,7 @@ class OllamaChatRequest(CustomBaseModel): class OllamaChatResponse(CustomBaseModel): """Chat response from Ollama API.""" + model: str created_at: str message: OllamaMessage @@ -31,6 +35,7 @@ class OllamaChatResponse(CustomBaseModel): class OllamaModelInfo(CustomBaseModel): """Model information from Ollama.""" + name: str modified_at: str size: int @@ -39,4 +44,5 @@ class OllamaModelInfo(CustomBaseModel): class OllamaModelsResponse(CustomBaseModel): """Response from Ollama models list endpoint.""" + models: list[OllamaModelInfo] diff --git a/src/responses/__init__.py b/src/responses/__init__.py index e0bd730..568adcb 100644 --- a/src/responses/__init__.py +++ b/src/responses/__init__.py @@ -9,12 +9,12 @@ This module implements the OpenAI Responses API format: """ from src.responses.schemas import ( + FunctionCallOutputItem, + MessageOutputItem, + OutputItem, + ReasoningOutputItem, Response, ResponseRequest, - OutputItem, - MessageOutputItem, - ReasoningOutputItem, - FunctionCallOutputItem, ) __all__ = [ diff --git a/src/responses/context.py b/src/responses/context.py index 773a0a8..a255004 100644 --- a/src/responses/context.py +++ b/src/responses/context.py @@ -6,6 +6,7 @@ Handles: - Context trimming to fit model limits - Reserve tokens for output generation """ + from typing import Any @@ -59,11 +60,7 @@ class ContextWindow: # Approximate: 4 characters per token return total_chars // 4 - async def trim_to_fit( - self, - items: list[Any], - reserve_tokens: int = 512 - ) -> list[Any]: + async def trim_to_fit(self, items: list[Any], reserve_tokens: int = 512) -> list[Any]: """ Trim items to fit within context window. @@ -102,11 +99,7 @@ class ContextWindow: return kept_items - async def fits_in_context( - self, - items: list[Any], - reserve_tokens: int = 512 - ) -> bool: + async def fits_in_context(self, items: list[Any], reserve_tokens: int = 512) -> bool: """ Check if items fit within context window. @@ -121,11 +114,7 @@ class ContextWindow: available_tokens = self.max_tokens - reserve_tokens return total_tokens <= available_tokens - async def get_usage_stats( - self, - items: list[Any], - reserve_tokens: int = 512 - ) -> dict: + async def get_usage_stats(self, items: list[Any], reserve_tokens: int = 512) -> dict: """ Get context window usage statistics. @@ -153,7 +142,7 @@ class ContextWindow: "reserved_tokens": reserve_tokens, "available_tokens": available_tokens, "usage_percent": round(usage_percent, 2), - "fits": total_tokens <= available_tokens + "fits": total_tokens <= available_tokens, } # ======================================================================== diff --git a/src/responses/history.py b/src/responses/history.py index 5f2980e..98cb053 100644 --- a/src/responses/history.py +++ b/src/responses/history.py @@ -6,8 +6,8 @@ Supports hybrid approach: - Optional conversation_id in metadata for server-side grouping - Server can augment with vector memories (future) """ + import hashlib -from typing import Dict, List from src.responses.schemas import Response, ResponseRequest @@ -36,7 +36,7 @@ class ConversationHistory: Args: max_turns: Maximum number of response turns to keep per conversation """ - self._conversations: Dict[str, List[Response]] = {} + self._conversations: dict[str, list[Response]] = {} self._max_turns = max_turns async def get_conversation_id(self, request: ResponseRequest) -> str: @@ -61,11 +61,7 @@ class ConversationHistory: first_msg = str(request.input[0]) if request.input else "" return hashlib.sha256(first_msg.encode()).hexdigest()[:16] - async def add_response( - self, - conversation_id: str, - response: Response - ) -> None: + async def add_response(self, conversation_id: str, response: Response) -> None: """ Add response to conversation history. @@ -81,7 +77,7 @@ class ConversationHistory: # Trim old turns to stay within limit await self._trim_history(conversation_id) - async def get_history(self, conversation_id: str) -> List[Response]: + async def get_history(self, conversation_id: str) -> list[Response]: """ Retrieve conversation history. @@ -125,20 +121,17 @@ class ConversationHistory: conversation_id: Conversation identifier """ if len(self._conversations[conversation_id]) > self._max_turns: - self._conversations[conversation_id] = ( - self._conversations[conversation_id][-self._max_turns:] - ) + self._conversations[conversation_id] = self._conversations[conversation_id][ + -self._max_turns : + ] # ======================================================================== # Future: Vector Memory Integration # ======================================================================== async def get_relevant_memories( - self, - conversation_id: str, - query: str, - limit: int = 5 - ) -> List[dict]: + self, conversation_id: str, query: str, limit: int = 5 + ) -> list[dict]: """ Retrieve relevant memories from vector store. diff --git a/src/responses/router.py b/src/responses/router.py index 54731b1..e6570fa 100644 --- a/src/responses/router.py +++ b/src/responses/router.py @@ -7,10 +7,10 @@ OpenAI-compatible /v1/responses endpoint with streaming support. from fastapi import APIRouter, HTTPException from sse_starlette.sse import EventSourceResponse -from src.responses import service -from src.responses.schemas import ResponseRequest, Response -from src.core.exceptions import ModelNotFoundError, AppException +from src.core.exceptions import AppException, ModelNotFoundError from src.core.logging_config import get_logger +from src.responses import service +from src.responses.schemas import Response, ResponseRequest logger = get_logger(__name__) @@ -58,14 +58,11 @@ async def create_response( if use_steward: logger.info("Streaming with Steward preprocessing for Tatlock request") from src.responses.streaming import StreamingCoordinator + coordinator = StreamingCoordinator() - return EventSourceResponse( - coordinator.stream_response_with_steward(request) - ) + return EventSourceResponse(coordinator.stream_response_with_steward(request)) else: - return EventSourceResponse( - service.create_response_stream(request) - ) + return EventSourceResponse(service.create_response_stream(request)) # Non-streaming response if use_steward: diff --git a/src/responses/schemas.py b/src/responses/schemas.py index bf3a7b6..5736b1b 100644 --- a/src/responses/schemas.py +++ b/src/responses/schemas.py @@ -8,18 +8,20 @@ OpenAI Responses API format with support for: - Streaming and non-streaming modes """ -from typing import Literal, Any +from typing import Any, Literal + from pydantic import Field, field_validator from src.core.models import CustomBaseModel - # ============================================================================ # Output Item Schemas (appear in response.output array) # ============================================================================ + class OutputTextContent(CustomBaseModel): """Text content in message output.""" + type: Literal["output_text"] = "output_text" text: str annotations: list[dict] = Field(default_factory=list) @@ -31,6 +33,7 @@ class MessageOutputItem(CustomBaseModel): Represents the assistant's final response message. """ + type: Literal["message"] = "message" id: str role: Literal["assistant"] = "assistant" @@ -45,6 +48,7 @@ class ReasoningOutputItem(CustomBaseModel): Represents the model's thinking/reasoning process. Displayed separately from the final answer. """ + type: Literal["reasoning"] = "reasoning" id: str summary: list[str] # List of reasoning steps @@ -57,6 +61,7 @@ class FunctionCallOutputItem(CustomBaseModel): Represents a tool/function that the model wants to execute. """ + type: Literal["function_call"] = "function_call" id: str name: str @@ -73,8 +78,10 @@ OutputItem = MessageOutputItem | ReasoningOutputItem | FunctionCallOutputItem # # Usage Tracking # ============================================================================ + class ResponseUsage(CustomBaseModel): """Token usage statistics for the response.""" + input_tokens: int output_tokens: int reasoning_tokens: int = 0 @@ -85,8 +92,10 @@ class ResponseUsage(CustomBaseModel): # Request Schema # ============================================================================ + class Tool(CustomBaseModel): """Tool/function definition.""" + name: str description: str parameters: dict[str, Any] @@ -94,6 +103,7 @@ class Tool(CustomBaseModel): class ReasoningConfig(CustomBaseModel): """Reasoning configuration.""" + effort: Literal["none", "minimal", "low", "medium", "high", "xhigh"] = "medium" summary: Literal["auto", "off"] = "auto" @@ -104,46 +114,25 @@ class ResponseRequest(CustomBaseModel): OpenAI Responses API format with optional extensions. """ + model: str = Field(description="Model ID to use") - input: list[dict] = Field( - description="Input messages or previous responses" - ) + input: list[dict] = Field(description="Input messages or previous responses") reasoning: dict[str, Any] | None = Field( - default=None, - description="Reasoning configuration: {effort: 'medium', summary: 'auto'}" - ) - tools: list[dict] | None = Field( - default=None, - description="Available tools/functions" + default=None, description="Reasoning configuration: {effort: 'medium', summary: 'auto'}" ) + tools: list[dict] | None = Field(default=None, description="Available tools/functions") metadata: dict[str, Any] | None = Field( - default=None, - description="Custom metadata (e.g., conversation_id for server-side tracking)" - ) - stream: bool = Field( - default=False, - description="Enable streaming mode" - ) - max_output_tokens: int | None = Field( - default=None, - description="Maximum tokens to generate" - ) - temperature: float = Field( - default=1.0, - ge=0.0, - le=2.0, - description="Sampling temperature" - ) - stop: list[str] | None = Field( - default=None, - description="Stop sequences" + default=None, description="Custom metadata (e.g., conversation_id for server-side tracking)" ) + stream: bool = Field(default=False, description="Enable streaming mode") + max_output_tokens: int | None = Field(default=None, description="Maximum tokens to generate") + temperature: float = Field(default=1.0, ge=0.0, le=2.0, description="Sampling temperature") + stop: list[str] | None = Field(default=None, description="Stop sequences") user: str | None = Field( - default=None, - description="Unique identifier for end-user (OpenAI standard)" + default=None, description="Unique identifier for end-user (OpenAI standard)" ) - @field_validator('reasoning') + @field_validator("reasoning") @classmethod def validate_reasoning(cls, v: dict[str, Any] | None) -> dict[str, Any] | None: """ @@ -154,21 +143,21 @@ class ResponseRequest(CustomBaseModel): - summary must be 'auto' or 'off' """ if v is not None: - if 'effort' in v: - allowed_efforts = ['none', 'minimal', 'low', 'medium', 'high', 'xhigh'] - if v['effort'] not in allowed_efforts: + if "effort" in v: + allowed_efforts = ["none", "minimal", "low", "medium", "high", "xhigh"] + if v["effort"] not in allowed_efforts: raise ValueError( f"reasoning.effort must be one of {allowed_efforts}, got '{v['effort']}'" ) - if 'summary' in v: - allowed_summaries = ['auto', 'off'] - if v['summary'] not in allowed_summaries: + if "summary" in v: + allowed_summaries = ["auto", "off"] + if v["summary"] not in allowed_summaries: raise ValueError( f"reasoning.summary must be one of {allowed_summaries}, got '{v['summary']}'" ) return v - @field_validator('max_output_tokens') + @field_validator("max_output_tokens") @classmethod def validate_max_output_tokens(cls, v: int | None) -> int | None: """ @@ -180,7 +169,7 @@ class ResponseRequest(CustomBaseModel): raise ValueError(f"max_output_tokens must be positive, got {v}") return v - @field_validator('stop') + @field_validator("stop") @classmethod def validate_stop_sequences(cls, v: list[str] | None) -> list[str] | None: """ @@ -203,20 +192,20 @@ class ResponseRequest(CustomBaseModel): # Response Schema # ============================================================================ + class Response(CustomBaseModel): """ Complete response object. Contains output array with reasoning, function calls, and messages. """ + id: str = Field(description="Unique response ID") object: Literal["response"] = "response" created_at: int = Field(description="Unix timestamp") model: str = Field(description="Model used") status: Literal["completed", "in_progress", "failed", "cancelled"] - output: list[OutputItem] = Field( - description="Output items (reasoning, function_call, message)" - ) + output: list[OutputItem] = Field(description="Output items (reasoning, function_call, message)") usage: ResponseUsage = Field(description="Token usage statistics") @@ -224,8 +213,10 @@ class Response(CustomBaseModel): # Error Schema # ============================================================================ + class ErrorDetail(CustomBaseModel): """Error detail object.""" + type: str message: str code: int | None = None @@ -233,4 +224,5 @@ class ErrorDetail(CustomBaseModel): class ErrorResponse(CustomBaseModel): """Error response format.""" + error: ErrorDetail diff --git a/src/responses/service.py b/src/responses/service.py index f1ebca3..353e9fb 100644 --- a/src/responses/service.py +++ b/src/responses/service.py @@ -52,9 +52,9 @@ def _extract_response_preview(response: Response) -> str: """Extract response preview text for tracing.""" if response.output: for item in response.output: - if hasattr(item, 'content'): + if hasattr(item, "content"): for content in item.content: - if hasattr(content, 'text'): + if hasattr(content, "text"): return content.text[:200] return "" @@ -79,10 +79,12 @@ async def _execute_single_delegation( result summary is a curated user-safe sentence. """ import time + start_time = time.time() if agent_name == "biographer": from src.agents.delegation import delegate_to_biographer + result = await delegate_to_biographer(task=task, context=context) duration = time.time() - start_time await tracker.track_call("delegate_to_biographer", duration) @@ -90,6 +92,7 @@ async def _execute_single_delegation( elif agent_name == "librarian": from src.agents.delegation import delegate_to_librarian + result = await delegate_to_librarian(task=task, context=context) duration = time.time() - start_time await tracker.track_call("delegate_to_librarian", duration) @@ -97,6 +100,7 @@ async def _execute_single_delegation( elif agent_name == "housekeeper": from src.agents.delegation import delegate_to_housekeeper + result = await delegate_to_housekeeper(task=task, context=context) duration = time.time() - start_time await tracker.track_call("delegate_to_housekeeper", duration) @@ -107,9 +111,7 @@ async def _execute_single_delegation( async def _handle_text_delegation( - response: str, - tracker: "ToolCallTracker", - conversation_id: str + response: str, tracker: "ToolCallTracker", conversation_id: str ) -> str: """ Handle text-based delegation fallback. @@ -173,8 +175,7 @@ async def _handle_text_delegation( conversation_id=conversation_id, ) tasks = [ - _execute_single_delegation(agent.lower(), task, tracker) - for agent, task in matches + _execute_single_delegation(agent.lower(), task, tracker) for agent, task in matches ] results = await asyncio.gather(*tasks, return_exceptions=True) @@ -188,10 +189,7 @@ async def _handle_text_delegation( got=len(results), conversation_id=conversation_id, ) - return ( - "I apologize, sir. I was unable to complete the " - "requested delegations." - ) + return "I apologize, sir. I was unable to complete the " "requested delegations." # Combine results (failures carry curated user-safe sentences) summaries = [] @@ -205,8 +203,7 @@ async def _handle_text_delegation( conversation_id=conversation_id, ) summaries.append( - f"**{agent_name}**: " - f"{get_think_message(agent_name, task, 'error')}" + f"**{agent_name}**: " f"{get_think_message(agent_name, task, 'error')}" ) else: _, output, _ = item @@ -226,9 +223,7 @@ async def _handle_text_delegation( conversation_id=conversation_id, ) try: - _, output, _ = await _execute_single_delegation( - agent_name, task, tracker - ) + _, output, _ = await _execute_single_delegation(agent_name, task, tracker) summaries.append(output) except Exception as e: logger.error( @@ -421,7 +416,7 @@ def _calculate_usage(input_messages: list[dict], output_items: list) -> Response elif isinstance(item, FunctionCallOutputItem): func_text = item.arguments output_tokens += len(func_text) // 4 - elif hasattr(item, 'type'): + elif hasattr(item, "type"): # Agent OutputItem objects (backward compatibility) if item.type == "reasoning": reasoning_text = " ".join(item.data.get("summary", [])) @@ -439,7 +434,7 @@ def _calculate_usage(input_messages: list[dict], output_items: list) -> Response input_tokens=input_tokens, output_tokens=output_tokens, reasoning_tokens=reasoning_tokens, - total_tokens=total_tokens + total_tokens=total_tokens, ) @@ -527,7 +522,7 @@ async def create_response(request: ResponseRequest) -> Response: model=request.model, status="completed", output=converted_items, - usage=usage + usage=usage, ) # Track conversation history (for analytics and future vector memory) @@ -640,12 +635,15 @@ async def create_response_with_steward(request: ResponseRequest) -> Response: # If Steward recommends ONLY delegation agents (biographer/librarian/housekeeper), # we still use two-phase but delegate directly in Phase 1 delegation_agents = {"biographer", "librarian", "housekeeper"} - delegation_only = all( - cap in delegation_agents - for cap in enriched.recommendation.recommended_capabilities - ) and enriched.recommendation.recommended_capabilities + delegation_only = ( + all( + cap in delegation_agents for cap in enriched.recommendation.recommended_capabilities + ) + and enriched.recommendation.recommended_capabilities + ) from src.agents.tatlock import TatlockAgent + tatlock = TatlockAgent() # Use enriched query (with location/timezone context) if available @@ -677,7 +675,9 @@ async def create_response_with_steward(request: ResponseRequest) -> Response: ) # Add text delegation results to expert_results if text_delegation_results != orchestration_results["raw_output"]: - orchestration_results["expert_results"]["text_delegation"] = text_delegation_results + orchestration_results["expert_results"]["text_delegation"] = ( + text_delegation_results + ) # Phase 2: Synthesize butler-toned response from all results tatlock_response = await tatlock.synthesize_from_results( @@ -694,23 +694,25 @@ async def create_response_with_steward(request: ResponseRequest) -> Response: # Add Steward reasoning as reasoning output if enriched.steward_reasoning: - output_items.append(ReasoningOutputItem( - id=f"rs_{generate_id()}", - summary=[enriched.steward_reasoning], - status="completed" - )) + output_items.append( + ReasoningOutputItem( + id=f"rs_{generate_id()}", + summary=[enriched.steward_reasoning], + status="completed", + ) + ) # Add Tatlock's message - output_items.append(MessageOutputItem( - id=f"msg_{generate_id()}", - role="assistant", - content=[OutputTextContent( - type="output_text", - text=tatlock_response, - annotations=[] - )], - status="completed" - )) + output_items.append( + MessageOutputItem( + id=f"msg_{generate_id()}", + role="assistant", + content=[ + OutputTextContent(type="output_text", text=tatlock_response, annotations=[]) + ], + status="completed", + ) + ) # Calculate usage (approximate) usage = _calculate_usage(request.input, output_items) @@ -721,7 +723,7 @@ async def create_response_with_steward(request: ResponseRequest) -> Response: model=request.model, status="completed", output=output_items, - usage=usage + usage=usage, ) # Track conversation history @@ -752,9 +754,7 @@ async def create_response_with_steward(request: ResponseRequest) -> Response: raise -async def create_response_stream( - request: ResponseRequest -) -> AsyncGenerator[dict, None]: +async def create_response_stream(request: ResponseRequest) -> AsyncGenerator[dict, None]: """ Create streaming response. @@ -774,10 +774,7 @@ async def create_response_stream( coordinator = StreamingCoordinator() async for event in coordinator.stream_response(request): - yield { - "event": event.event, - "data": event.model_dump_json() - } + yield {"event": event.event, "data": event.model_dump_json()} async def get_conversation_history(conversation_id: str) -> list[Response]: @@ -802,7 +799,7 @@ async def get_conversation_stats() -> dict: """ return { "total_conversations": await _conversation_history.get_conversation_count(), - "max_turns_per_conversation": _conversation_history._max_turns + "max_turns_per_conversation": _conversation_history._max_turns, } @@ -843,23 +840,29 @@ def _convert_output_items(items: list) -> list: for item in items: if item.type == "message": - converted.append(MessageOutputItem( - id=item.id, - content=[OutputTextContent(**c) for c in item.data["content"]], - status=item.data.get("status", "completed") - )) + converted.append( + MessageOutputItem( + id=item.id, + content=[OutputTextContent(**c) for c in item.data["content"]], + status=item.data.get("status", "completed"), + ) + ) elif item.type == "reasoning": - converted.append(ReasoningOutputItem( - id=item.id, - summary=item.data["summary"], - status=item.data.get("status", "completed") - )) + converted.append( + ReasoningOutputItem( + id=item.id, + summary=item.data["summary"], + status=item.data.get("status", "completed"), + ) + ) elif item.type == "function_call": - converted.append(FunctionCallOutputItem( - id=item.id, - name=item.data["name"], - arguments=item.data["arguments"], - status=item.data.get("status", "completed") - )) + converted.append( + FunctionCallOutputItem( + id=item.id, + name=item.data["name"], + arguments=item.data["arguments"], + status=item.data.get("status", "completed"), + ) + ) return converted diff --git a/src/responses/streaming.py b/src/responses/streaming.py index 0db4ec8..85d62e2 100644 --- a/src/responses/streaming.py +++ b/src/responses/streaming.py @@ -30,8 +30,10 @@ logger = get_logger(__name__) # Stream Event Types # ============================================================================ + class StreamEventType(str, Enum): """Streaming event types for Responses API.""" + REASONING_SUMMARY_DELTA = "response.reasoning_summary_text.delta" REASONING_SUMMARY_DONE = "response.reasoning_summary_text.done" OUTPUT_TEXT_DELTA = "response.output_text.delta" @@ -46,30 +48,38 @@ class StreamEventType(str, Enum): # Stream Event Schemas # ============================================================================ + class ReasoningSummaryDelta(CustomBaseModel): """Reasoning summary text delta event.""" - event: Literal[StreamEventType.REASONING_SUMMARY_DELTA] = StreamEventType.REASONING_SUMMARY_DELTA + + event: Literal[StreamEventType.REASONING_SUMMARY_DELTA] = ( + StreamEventType.REASONING_SUMMARY_DELTA + ) delta: str class ReasoningSummaryDone(CustomBaseModel): """Reasoning summary completion event.""" + event: Literal[StreamEventType.REASONING_SUMMARY_DONE] = StreamEventType.REASONING_SUMMARY_DONE class OutputTextDelta(CustomBaseModel): """Output text delta event.""" + event: Literal[StreamEventType.OUTPUT_TEXT_DELTA] = StreamEventType.OUTPUT_TEXT_DELTA delta: str class OutputTextDone(CustomBaseModel): """Output text completion event.""" + event: Literal[StreamEventType.OUTPUT_TEXT_DONE] = StreamEventType.OUTPUT_TEXT_DONE class FunctionCallDelta(CustomBaseModel): """Function call arguments delta event.""" + event: Literal[StreamEventType.FUNCTION_CALL_DELTA] = StreamEventType.FUNCTION_CALL_DELTA delta: str name: str | None = None # Only in first chunk @@ -77,31 +87,34 @@ class FunctionCallDelta(CustomBaseModel): class FunctionCallDone(CustomBaseModel): """Function call completion event.""" + event: Literal[StreamEventType.FUNCTION_CALL_DONE] = StreamEventType.FUNCTION_CALL_DONE class ResponseDone(CustomBaseModel): """Response completion event with full response.""" + event: Literal[StreamEventType.RESPONSE_DONE] = StreamEventType.RESPONSE_DONE response: Response class ErrorEvent(CustomBaseModel): """Error event.""" + event: Literal[StreamEventType.ERROR] = StreamEventType.ERROR error: dict # Union type for all stream events StreamEvent = ( - ReasoningSummaryDelta | - ReasoningSummaryDone | - OutputTextDelta | - OutputTextDone | - FunctionCallDelta | - FunctionCallDone | - ResponseDone | - ErrorEvent + ReasoningSummaryDelta + | ReasoningSummaryDone + | OutputTextDelta + | OutputTextDone + | FunctionCallDelta + | FunctionCallDone + | ResponseDone + | ErrorEvent ) @@ -109,6 +122,7 @@ StreamEvent = ( # Streaming Coordinator # ============================================================================ + class StreamingCoordinator: """ Coordinates streaming from agents to SSE format. @@ -123,7 +137,7 @@ class StreamingCoordinator: async def stream_response_with_steward( self, - request: "ResponseRequest" # type: ignore # Forward reference + request: "ResponseRequest", # type: ignore # Forward reference ) -> AsyncGenerator[StreamEvent, None]: """ Stream response with Steward preprocessing and two-phase Tatlock execution. @@ -181,10 +195,13 @@ class StreamingCoordinator: # Check if direct delegation is recommended delegation_agents = {"biographer", "librarian", "housekeeper"} - delegation_only = all( - cap in delegation_agents - for cap in enriched.recommendation.recommended_capabilities - ) and enriched.recommendation.recommended_capabilities + delegation_only = ( + all( + cap in delegation_agents + for cap in enriched.recommendation.recommended_capabilities + ) + and enriched.recommendation.recommended_capabilities + ) tatlock = TatlockAgent() @@ -222,7 +239,7 @@ class StreamingCoordinator: # Stream the synthesized response chunk_size = 50 for i in range(0, len(tatlock_response), chunk_size): - yield OutputTextDelta(delta=tatlock_response[i:i + chunk_size]) + yield OutputTextDelta(delta=tatlock_response[i : i + chunk_size]) await asyncio.sleep(0.02) yield OutputTextDone() @@ -231,12 +248,10 @@ class StreamingCoordinator: message_item = MessageOutputItem( id=f"msg_{generate_id()}", role="assistant", - content=[OutputTextContent( - type="output_text", - text=tatlock_response, - annotations=[] - )], - status="completed" + content=[ + OutputTextContent(type="output_text", text=tatlock_response, annotations=[]) + ], + status="completed", ) output_items.append(message_item) @@ -252,7 +267,7 @@ class StreamingCoordinator: model=request.model, status="completed", output=output_items, - usage=usage + usage=usage, ) # Track conversation history @@ -319,17 +334,11 @@ class StreamingCoordinator: try: # Execute delegation if agent == "librarian": - result = await delegate_to_librarian( - task=user_message, context=context - ) + result = await delegate_to_librarian(task=user_message, context=context) elif agent == "biographer": - result = await delegate_to_biographer( - task=user_message, context=context - ) + result = await delegate_to_biographer(task=user_message, context=context) elif agent == "housekeeper": - result = await delegate_to_housekeeper( - task=user_message, context=context - ) + result = await delegate_to_housekeeper(task=user_message, context=context) else: result = None @@ -366,17 +375,19 @@ class StreamingCoordinator: yield ReasoningSummaryDone() if results is not None: - results.update({ - "tools_called": tools_called, - "expert_results": expert_results, - "tool_outputs": {}, - "raw_output": "", - "think_messages": think_messages, - }) + results.update( + { + "tools_called": tools_called, + "expert_results": expert_results, + "tool_outputs": {}, + "raw_output": "", + "think_messages": think_messages, + } + ) async def stream_response( self, - request: "ResponseRequest" # type: ignore # Forward reference + request: "ResponseRequest", # type: ignore # Forward reference ) -> AsyncGenerator[StreamEvent, None]: """ Coordinate streaming from agent to SSE events. @@ -435,18 +446,13 @@ class StreamingCoordinator: elif item.type == "function_call": # Stream function call arguments # First chunk includes name - yield FunctionCallDelta( - name=item.data["name"], - delta="" - ) + yield FunctionCallDelta(name=item.data["name"], delta="") # Stream arguments in chunks args = item.data["arguments"] chunk_size = 20 for i in range(0, len(args), chunk_size): - yield FunctionCallDelta( - delta=args[i:i+chunk_size] - ) + yield FunctionCallDelta(delta=args[i : i + chunk_size]) await asyncio.sleep(0.03) yield FunctionCallDone() @@ -458,7 +464,7 @@ class StreamingCoordinator: # Only stream the NEW text (delta) since last update if current_text.startswith(last_message_text): # Extract only the new portion - delta_text = current_text[len(last_message_text):] + delta_text = current_text[len(last_message_text) :] if delta_text: # Stream the delta text in chunks while preserving formatting @@ -466,17 +472,16 @@ class StreamingCoordinator: chunk_size = 50 # characters per chunk for i in range(0, len(delta_text), chunk_size): - chunk = delta_text[i:i+chunk_size] + chunk = delta_text[i : i + chunk_size] # Check stop sequences on full accumulated text stop_found, text_before_stop = self._check_stop_sequence( - current_text, - request.stop + current_text, request.stop ) if stop_found: # Only emit remaining delta before stop - remaining = text_before_stop[len(last_message_text):] + remaining = text_before_stop[len(last_message_text) :] if remaining: yield OutputTextDelta(delta=remaining) yield OutputTextDone() @@ -508,11 +513,12 @@ class StreamingCoordinator: model=request.model, status="completed", output=self._convert_output_items(output_items), - usage=usage + usage=usage, ) # Track conversation history (import here to avoid circular dependency) from src.responses.service import _conversation_history + conversation_id = await _conversation_history.get_conversation_id(request) await _conversation_history.add_response(conversation_id, final_response) @@ -534,24 +540,30 @@ class StreamingCoordinator: converted = [] for item in items: if item.type == "message": - converted.append(MessageOutputItem( - id=item.id, - content=[OutputTextContent(**c) for c in item.data["content"]], - status=item.data.get("status", "completed") - )) + converted.append( + MessageOutputItem( + id=item.id, + content=[OutputTextContent(**c) for c in item.data["content"]], + status=item.data.get("status", "completed"), + ) + ) elif item.type == "reasoning": - converted.append(ReasoningOutputItem( - id=item.id, - summary=item.data["summary"], - status=item.data.get("status", "completed") - )) + converted.append( + ReasoningOutputItem( + id=item.id, + summary=item.data["summary"], + status=item.data.get("status", "completed"), + ) + ) elif item.type == "function_call": - converted.append(FunctionCallOutputItem( - id=item.id, - name=item.data["name"], - arguments=item.data["arguments"], - status=item.data.get("status", "completed") - )) + converted.append( + FunctionCallOutputItem( + id=item.id, + name=item.data["name"], + arguments=item.data["arguments"], + status=item.data.get("status", "completed"), + ) + ) return converted @@ -576,18 +588,10 @@ class StreamingCoordinator: error_type = "internal_error" code = 500 - return ErrorEvent( - error={ - "type": error_type, - "message": str(error), - "code": code - } - ) + return ErrorEvent(error={"type": error_type, "message": str(error), "code": code}) def _check_stop_sequence( - self, - accumulated_text: str, - stop_sequences: list[str] | None + self, accumulated_text: str, stop_sequences: list[str] | None ) -> tuple[bool, str]: """ Check if any stop sequence is encountered. @@ -624,11 +628,7 @@ class StreamingCoordinator: """ return len(text) // 4 - def _check_max_tokens( - self, - current_tokens: int, - max_tokens: int | None - ) -> bool: + def _check_max_tokens(self, current_tokens: int, max_tokens: int | None) -> bool: """ Check if max tokens limit reached. diff --git a/tests/agents/biographer/test_capability.py b/tests/agents/biographer/test_capability.py index 3e68882..b635604 100644 --- a/tests/agents/biographer/test_capability.py +++ b/tests/agents/biographer/test_capability.py @@ -2,9 +2,10 @@ Tests for Biographer capability registration. """ -import pytest from unittest.mock import MagicMock, patch +import pytest + from src.agents.biographer.capability import ( BIOGRAPHER_CAPABILITY, get_biographer_capability, @@ -73,9 +74,7 @@ class TestBiographerRegistration: "src.agents.biographer.capability.get_household_registry", return_value=mock_registry, ): - with patch( - "src.agents.biographer.capability.get_biographer_agent" - ) as mock_get_agent: + with patch("src.agents.biographer.capability.get_biographer_agent") as mock_get_agent: mock_agent = MagicMock() mock_get_agent.return_value = mock_agent diff --git a/tests/agents/housekeeper/test_capability.py b/tests/agents/housekeeper/test_capability.py index c170572..1f43d0e 100644 --- a/tests/agents/housekeeper/test_capability.py +++ b/tests/agents/housekeeper/test_capability.py @@ -2,9 +2,10 @@ Tests for Housekeeper capability registration. """ -import pytest from unittest.mock import MagicMock, patch +import pytest + from src.agents.housekeeper.capability import ( HOUSEKEEPER_CAPABILITY, get_housekeeper_capability, @@ -74,9 +75,7 @@ class TestHousekeeperRegistration: "src.agents.housekeeper.capability.get_household_registry", return_value=mock_registry, ): - with patch( - "src.agents.housekeeper.capability.get_housekeeper_agent" - ) as mock_get_agent: + with patch("src.agents.housekeeper.capability.get_housekeeper_agent") as mock_get_agent: mock_agent = MagicMock() mock_get_agent.return_value = mock_agent diff --git a/tests/agents/housekeeper/test_client.py b/tests/agents/housekeeper/test_client.py index 9b73bc6..70d0a5d 100644 --- a/tests/agents/housekeeper/test_client.py +++ b/tests/agents/housekeeper/test_client.py @@ -2,9 +2,10 @@ Tests for the Core-API HTTP client. """ -import pytest from unittest.mock import AsyncMock, MagicMock + import httpx +import pytest from src.agents.housekeeper.client import ( Area, diff --git a/tests/agents/librarian/test_capability.py b/tests/agents/librarian/test_capability.py index 2c7495f..a3724c1 100644 --- a/tests/agents/librarian/test_capability.py +++ b/tests/agents/librarian/test_capability.py @@ -2,9 +2,10 @@ Tests for Librarian capability registration. """ -import pytest from unittest.mock import MagicMock, patch +import pytest + from src.agents.librarian.capability import ( LIBRARIAN_CAPABILITY, get_librarian_capability, @@ -67,9 +68,7 @@ class TestLibrarianRegistration: "src.agents.librarian.capability.get_household_registry", return_value=mock_registry, ): - with patch( - "src.agents.librarian.capability.get_librarian_agent" - ) as mock_get_agent: + with patch("src.agents.librarian.capability.get_librarian_agent") as mock_get_agent: mock_agent = MagicMock() mock_get_agent.return_value = mock_agent diff --git a/tests/agents/librarian/test_client.py b/tests/agents/librarian/test_client.py index 5d82c9c..bebeee7 100644 --- a/tests/agents/librarian/test_client.py +++ b/tests/agents/librarian/test_client.py @@ -137,9 +137,7 @@ class TestHybridSearch: assert "docker" in result.keywords @pytest.mark.asyncio - async def test_hybrid_search_empty_results( - self, client_with_mock, mock_httpx_client - ): + async def test_hybrid_search_empty_results(self, client_with_mock, mock_httpx_client): """Test hybrid search with no results.""" mock_response = MagicMock() mock_response.json.return_value = { @@ -445,9 +443,7 @@ class TestUpdateWikiPage: assert page.tags == ["projects", "devops"] @pytest.mark.asyncio - async def test_update_wiki_page_multiple_fields( - self, client_with_mock, mock_httpx_client - ): + async def test_update_wiki_page_multiple_fields(self, client_with_mock, mock_httpx_client): """Test updating multiple fields at once.""" mock_response = MagicMock() mock_response.json.return_value = { @@ -622,15 +618,11 @@ class TestExplicitUserContract: @pytest.mark.asyncio @pytest.mark.parametrize("method_name,kwargs", TENANT_SCOPED_METHODS) - async def test_user_from_context_is_sent_on_the_wire( - self, method_name, kwargs - ): + async def test_user_from_context_is_sent_on_the_wire(self, method_name, kwargs): """With no explicit user, the context user is resolved and sent.""" client, mock_httpx = self._wire_client() - with patch( - "src.agents.librarian.client.get_user", return_value="llm_tester" - ): + with patch("src.agents.librarian.client.get_user", return_value="llm_tester"): await getattr(client, method_name)(**kwargs) assert self._sent_user(mock_httpx) == "llm_tester" @@ -647,9 +639,7 @@ class TestExplicitUserContract: @pytest.mark.asyncio @pytest.mark.parametrize("method_name,kwargs", TENANT_SCOPED_METHODS) - async def test_empty_context_user_fails_before_any_request( - self, method_name, kwargs - ): + async def test_empty_context_user_fails_before_any_request(self, method_name, kwargs): """An empty resolved user raises before any bytes hit the wire.""" client, mock_httpx = self._wire_client() @@ -698,9 +688,7 @@ class TestExplicitUserContract: from src.core import config as config_module from src.core.config import Environment - monkeypatch.setattr( - config_module.config, "ENVIRONMENT", Environment.DEVELOPMENT - ) + monkeypatch.setattr(config_module.config, "ENVIRONMENT", Environment.DEVELOPMENT) client, mock_httpx = self._wire_client() await getattr(client, method_name)(user=explicit_user, **kwargs) @@ -708,16 +696,12 @@ class TestExplicitUserContract: assert self._sent_user(mock_httpx) == "llm_tester" @pytest.mark.asyncio - async def test_explicit_production_tenant_passes_through_in_prod( - self, monkeypatch - ): + async def test_explicit_production_tenant_passes_through_in_prod(self, monkeypatch): """In production the production tenant is sent unchanged.""" from src.core import config as config_module from src.core.config import Environment - monkeypatch.setattr( - config_module.config, "ENVIRONMENT", Environment.PRODUCTION - ) + monkeypatch.setattr(config_module.config, "ENVIRONMENT", Environment.PRODUCTION) client, mock_httpx = self._wire_client() await client.hybrid_search("q", user="jpmschweitzer") diff --git a/tests/agents/librarian/test_client_retry.py b/tests/agents/librarian/test_client_retry.py index 23035fb..ded0d37 100644 --- a/tests/agents/librarian/test_client_retry.py +++ b/tests/agents/librarian/test_client_retry.py @@ -91,9 +91,7 @@ class TestBoundedRetries: mock_httpx.post.side_effect = httpx.ConnectError("Connection refused") with pytest.raises(httpx.ConnectError): - await client_with_mock.create_wiki_page( - title="T", path="/t", content="c", user="u" - ) + await client_with_mock.create_wiki_page(title="T", path="/t", content="c", user="u") assert mock_httpx.post.call_count == 1 @@ -103,9 +101,7 @@ class TestBoundedRetries: mock_httpx.post.side_effect = httpx.ConnectError("Connection refused") with pytest.raises(httpx.ConnectError): - await client_with_mock.smart_create_wiki_page( - topic="T", tags=["x"], user="u" - ) + await client_with_mock.smart_create_wiki_page(topic="T", tags=["x"], user="u") assert mock_httpx.post.call_count == 1 @@ -184,9 +180,7 @@ class TestModelRetryEscalation: @pytest.mark.asyncio async def test_read_tool_raises_model_retry_on_transport_error(self): mock_client = AsyncMock() - mock_client.hybrid_search.side_effect = httpx.ConnectError( - "Connection refused" - ) + mock_client.hybrid_search.side_effect = httpx.ConnectError("Connection refused") with self._patched_client(mock_client): with pytest.raises(ModelRetry): @@ -219,14 +213,10 @@ class TestModelRetryEscalation: @pytest.mark.asyncio async def test_write_tool_never_raises_model_retry(self): mock_client = AsyncMock() - mock_client.create_wiki_page.side_effect = httpx.ConnectError( - "Connection refused" - ) + mock_client.create_wiki_page.side_effect = httpx.ConnectError("Connection refused") with self._patched_client(mock_client): - result = await create_wiki_page( - title="T", path="/t", content="c", tags=["x"] - ) + result = await create_wiki_page(title="T", path="/t", content="c", tags=["x"]) assert "unable" in result assert "Connection refused" not in result diff --git a/tests/agents/librarian/test_hybrid_contract.py b/tests/agents/librarian/test_hybrid_contract.py index df019be..a4068b7 100644 --- a/tests/agents/librarian/test_hybrid_contract.py +++ b/tests/agents/librarian/test_hybrid_contract.py @@ -83,9 +83,7 @@ class TestHybridRAGContract: assert source in SOURCE_ICONS, f"no icon for source '{source}'" @pytest.mark.asyncio - async def test_context_maps_to_formatted_context( - self, client_with_recorded_response - ): + async def test_context_maps_to_formatted_context(self, client_with_recorded_response): """Top-level 'context' field maps to formatted_context.""" response = await client_with_recorded_response.hybrid_search( "home server infrastructure", user="testuser" @@ -94,9 +92,7 @@ class TestHybridRAGContract: assert response.formatted_context != "" @pytest.mark.asyncio - async def test_keywords_and_synonyms_from_dict( - self, client_with_recorded_response - ): + async def test_keywords_and_synonyms_from_dict(self, client_with_recorded_response): """keywords is a dict: core_keywords + nested synonyms map.""" response = await client_with_recorded_response.hybrid_search( "home server infrastructure", user="testuser" @@ -124,9 +120,7 @@ class TestHybridRAGContract: assert len(response.related_dossiers) == len(set(response.related_dossiers)) @pytest.mark.asyncio - async def test_payload_never_sends_zero_limits( - self, client_with_recorded_response - ): + async def test_payload_never_sends_zero_limits(self, client_with_recorded_response): """The live service 422s on limits < 1; disabled legs use enable_* flags.""" await client_with_recorded_response.hybrid_search( "home server infrastructure", @@ -151,24 +145,18 @@ class TestHybridRAGContract: assert config["enable_volatile"] is False @pytest.mark.asyncio - async def test_user_always_sent_as_query_param( - self, client_with_recorded_response - ): + async def test_user_always_sent_as_query_param(self, client_with_recorded_response): """The tenant is always sent explicitly - library-desk is removing its server-side default, so a missing user would 422.""" await client_with_recorded_response.hybrid_search( "home server infrastructure", user="testuser" ) - params = client_with_recorded_response._client.post.call_args.kwargs[ - "params" - ] + params = client_with_recorded_response._client.post.call_args.kwargs["params"] assert params["user"] == "testuser" @pytest.mark.asyncio - async def test_source_counts_and_timing_parsed( - self, client_with_recorded_response - ): + async def test_source_counts_and_timing_parsed(self, client_with_recorded_response): """source_counts and timing map into the response model.""" response = await client_with_recorded_response.hybrid_search( "home server infrastructure", user="testuser" @@ -178,9 +166,7 @@ class TestHybridRAGContract: assert response.timing.get("total_ms", 0) > 0 @pytest.mark.asyncio - async def test_source_status_absent_is_tolerated( - self, client_with_recorded_response - ): + async def test_source_status_absent_is_tolerated(self, client_with_recorded_response): """Recorded response predates source_status/degraded - defaults apply.""" response = await client_with_recorded_response.hybrid_search( "home server infrastructure", user="testuser" @@ -218,7 +204,9 @@ class TestHybridRAGContract: assert response.source_status["volatile"] == "disabled" @pytest.mark.asyncio - async def test_tool_renders_no_unknown_results(self, client_with_recorded_response, monkeypatch): + async def test_tool_renders_no_unknown_results( + self, client_with_recorded_response, monkeypatch + ): """The hybrid_search tool renders real sources and non-zero scores.""" class _Factory: @@ -231,9 +219,7 @@ class TestHybridRAGContract: async def __aexit__(self, *args): return None - monkeypatch.setattr( - "src.agents.librarian.tools.LibraryDeskClient", _Factory() - ) + monkeypatch.setattr("src.agents.librarian.tools.LibraryDeskClient", _Factory()) output = await hybrid_search("home server infrastructure") @@ -278,9 +264,7 @@ class TestCoverageNote: def test_wiki_leg_absence_is_not_degradation(self): """vector/graph missing from top-N counts is healthy ranking, not outage.""" - response = self._response( - source_counts={"web": 2, "documents": 1, "volatile": 1} - ) + response = self._response(source_counts={"web": 2, "documents": 1, "volatile": 1}) note = _coverage_note( response, include_web=True, include_documents=True, include_volatile=True ) diff --git a/tests/agents/librarian/test_tool_schemas.py b/tests/agents/librarian/test_tool_schemas.py index 8e13a70..12cd865 100644 --- a/tests/agents/librarian/test_tool_schemas.py +++ b/tests/agents/librarian/test_tool_schemas.py @@ -35,9 +35,7 @@ def _nullable_anyof_paths(schema: object, path: str = "") -> list[str]: class TestLibrarianToolSchemas: """All registered librarian tools emit Ollama-safe parameter schemas.""" - @pytest.mark.parametrize( - "tool_func", LIBRARIAN_TOOLS, ids=lambda f: f.__name__ - ) + @pytest.mark.parametrize("tool_func", LIBRARIAN_TOOLS, ids=lambda f: f.__name__) def test_no_nullable_anyof_in_schema(self, tool_func): schema = Tool(tool_func).function_schema.json_schema diff --git a/tests/agents/librarian/test_tools.py b/tests/agents/librarian/test_tools.py index 516e6bd..1b362cb 100644 --- a/tests/agents/librarian/test_tools.py +++ b/tests/agents/librarian/test_tools.py @@ -36,6 +36,7 @@ def mock_client(): # Web Search Tests # ============================================================================ + @pytest.mark.unit class TestSearchWeb: """Tests for search_web tool.""" @@ -68,9 +69,7 @@ class TestSearchWeb: ) mock_client.search_web.return_value = mock_response - with patch( - "src.agents.librarian.tools.LibraryDeskClient" - ) as mock_client_class: + with patch("src.agents.librarian.tools.LibraryDeskClient") as mock_client_class: mock_client_class.return_value.__aenter__.return_value = mock_client mock_client_class.return_value.__aexit__.return_value = None @@ -94,9 +93,7 @@ class TestSearchWeb: ) mock_client.search_web.return_value = mock_response - with patch( - "src.agents.librarian.tools.LibraryDeskClient" - ) as mock_client_class: + with patch("src.agents.librarian.tools.LibraryDeskClient") as mock_client_class: mock_client_class.return_value.__aenter__.return_value = mock_client mock_client_class.return_value.__aexit__.return_value = None @@ -111,9 +108,7 @@ class TestSearchWeb: "Connection failed to http://internal-host:8089" ) - with patch( - "src.agents.librarian.tools.LibraryDeskClient" - ) as mock_client_class: + with patch("src.agents.librarian.tools.LibraryDeskClient") as mock_client_class: mock_client_class.return_value.__aenter__.return_value = mock_client mock_client_class.return_value.__aexit__.return_value = None @@ -144,9 +139,7 @@ class TestSearchWeb: ) mock_client.search_web.return_value = mock_response - with patch( - "src.agents.librarian.tools.LibraryDeskClient" - ) as mock_client_class: + with patch("src.agents.librarian.tools.LibraryDeskClient") as mock_client_class: mock_client_class.return_value.__aenter__.return_value = mock_client mock_client_class.return_value.__aexit__.return_value = None @@ -164,6 +157,7 @@ class TestSearchWeb: # Read URL Tests # ============================================================================ + @pytest.mark.unit class TestReadUrl: """Tests for read_url tool.""" @@ -182,9 +176,7 @@ class TestReadUrl: ) mock_client.extract_content.return_value = mock_result - with patch( - "src.agents.librarian.tools.LibraryDeskClient" - ) as mock_client_class: + with patch("src.agents.librarian.tools.LibraryDeskClient") as mock_client_class: mock_client_class.return_value.__aenter__.return_value = mock_client mock_client_class.return_value.__aexit__.return_value = None @@ -205,9 +197,7 @@ class TestReadUrl: ) mock_client.extract_content.return_value = mock_result - with patch( - "src.agents.librarian.tools.LibraryDeskClient" - ) as mock_client_class: + with patch("src.agents.librarian.tools.LibraryDeskClient") as mock_client_class: mock_client_class.return_value.__aenter__.return_value = mock_client mock_client_class.return_value.__aexit__.return_value = None @@ -227,9 +217,7 @@ class TestReadUrl: ) mock_client.extract_content.return_value = mock_result - with patch( - "src.agents.librarian.tools.LibraryDeskClient" - ) as mock_client_class: + with patch("src.agents.librarian.tools.LibraryDeskClient") as mock_client_class: mock_client_class.return_value.__aenter__.return_value = mock_client mock_client_class.return_value.__aexit__.return_value = None @@ -246,6 +234,7 @@ class TestReadUrl: # Batch URL Tests # ============================================================================ + @pytest.mark.unit class TestReadUrlsBatch: """Tests for read_urls_batch tool.""" @@ -275,16 +264,16 @@ class TestReadUrlsBatch: ) mock_client.extract_content_batch.return_value = mock_response - with patch( - "src.agents.librarian.tools.LibraryDeskClient" - ) as mock_client_class: + with patch("src.agents.librarian.tools.LibraryDeskClient") as mock_client_class: mock_client_class.return_value.__aenter__.return_value = mock_client mock_client_class.return_value.__aexit__.return_value = None - result = await read_urls_batch([ - "https://example.com/1", - "https://example.com/2", - ]) + result = await read_urls_batch( + [ + "https://example.com/1", + "https://example.com/2", + ] + ) assert "Article 1" in result assert "Article 2" in result @@ -314,16 +303,16 @@ class TestReadUrlsBatch: ) mock_client.extract_content_batch.return_value = mock_response - with patch( - "src.agents.librarian.tools.LibraryDeskClient" - ) as mock_client_class: + with patch("src.agents.librarian.tools.LibraryDeskClient") as mock_client_class: mock_client_class.return_value.__aenter__.return_value = mock_client mock_client_class.return_value.__aexit__.return_value = None - result = await read_urls_batch([ - "https://example.com/good", - "https://example.com/bad", - ]) + result = await read_urls_batch( + [ + "https://example.com/good", + "https://example.com/bad", + ] + ) # Should contain successful result assert "Good Article" in result @@ -336,6 +325,7 @@ class TestReadUrlsBatch: # Response Model Tests # ============================================================================ + @pytest.mark.unit class TestWebSearchModels: """Tests for web search response models.""" @@ -436,6 +426,7 @@ class TestWebSearchModels: # Wiki Update Tests (tag sentinel behavior) # ============================================================================ + @pytest.mark.unit class TestUpdateWikiPageTagSentinels: """Empty list leaves tags unchanged; the clear sentinel empties them.""" diff --git a/tests/agents/steward/test_steward_schemas.py b/tests/agents/steward/test_steward_schemas.py index 901d8a1..e106a70 100644 --- a/tests/agents/steward/test_steward_schemas.py +++ b/tests/agents/steward/test_steward_schemas.py @@ -3,7 +3,6 @@ Tests for Steward schemas. Tests the structured output models for conversation context and recommendations. """ -import pytest from src.agents.steward.schemas import ConversationContext, StewardRecommendation @@ -24,7 +23,7 @@ class TestConversationContext: context = ConversationContext( has_previous_context=True, relevant_turns=[0, 2, 4], - context_summary="User discussed weather in turns 0 and 2" + context_summary="User discussed weather in turns 0 and 2", ) assert context.has_previous_context is True @@ -89,7 +88,7 @@ class TestStewardRecommendation: context = ConversationContext( has_previous_context=True, relevant_turns=[1, 3], - context_summary="User asked about calculation in turn 1, now wants explanation" + context_summary="User asked about calculation in turn 1, now wants explanation", ) rec = StewardRecommendation( @@ -122,7 +121,7 @@ class TestStewardRecommendation: context = ConversationContext( has_previous_context=True, relevant_turns=[0], - context_summary="Previous calculation mentioned" + context_summary="Previous calculation mentioned", ) rec = StewardRecommendation( diff --git a/tests/agents/steward/test_steward_service.py b/tests/agents/steward/test_steward_service.py index e0516ce..5862165 100644 --- a/tests/agents/steward/test_steward_service.py +++ b/tests/agents/steward/test_steward_service.py @@ -3,6 +3,7 @@ Tests for Steward service layer. Tests request analysis, logging, and benchmarking integration. """ + from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -31,7 +32,9 @@ class TestAnalyzeRequest: """Test analyzing a simple greeting.""" # Mock the Steward agent's analyze method (plain text approach) mock_agent = MagicMock() - mock_agent.analyze = AsyncMock(return_value="Simple greeting requires no tools. This is a simple request.") + mock_agent.analyze = AsyncMock( + return_value="Simple greeting requires no tools. This is a simple request." + ) with patch("src.agents.steward.service.get_steward_agent", return_value=mock_agent): result = await analyze_request( @@ -159,7 +162,7 @@ class TestFormatStewardNote: context = ConversationContext( has_previous_context=True, relevant_turns=[0, 1], - context_summary="Previous discussion about calculations" + context_summary="Previous discussion about calculations", ) rec = StewardRecommendation( @@ -205,9 +208,7 @@ class TestBuildEnrichedQuery: def test_enrichment_adds_location(self): """Test location is appended for weather queries.""" query = "What's the weather?" - memory_context = { - "profile": {"location": "Amsterdam", "timezone": "Europe/Amsterdam"} - } + memory_context = {"profile": {"location": "Amsterdam", "timezone": "Europe/Amsterdam"}} result = _build_enriched_query(query, memory_context) @@ -218,9 +219,7 @@ class TestBuildEnrichedQuery: def test_no_location_when_specified(self): """Test location is not appended when already specified.""" query = "What's the weather in London?" - memory_context = { - "profile": {"location": "Amsterdam"} - } + memory_context = {"profile": {"location": "Amsterdam"}} result = _build_enriched_query(query, memory_context) @@ -230,9 +229,7 @@ class TestBuildEnrichedQuery: def test_enrichment_adds_timezone(self): """Test timezone is appended for time queries.""" query = "What time is it?" - memory_context = { - "profile": {"timezone": "Europe/Amsterdam"} - } + memory_context = {"profile": {"timezone": "Europe/Amsterdam"}} result = _build_enriched_query(query, memory_context) @@ -241,9 +238,7 @@ class TestBuildEnrichedQuery: def test_no_timezone_when_specified(self): """Test timezone is not appended when already specified.""" query = "What time is it in UTC?" - memory_context = { - "profile": {"timezone": "Europe/Amsterdam"} - } + memory_context = {"profile": {"timezone": "Europe/Amsterdam"}} result = _build_enriched_query(query, memory_context) @@ -254,7 +249,7 @@ class TestBuildEnrichedQuery: query = "What's the weather?" memory_context = { "profile": {"location": "Amsterdam"}, - "preferences": {"temperature_unit": "celsius"} + "preferences": {"temperature_unit": "celsius"}, } result = _build_enriched_query(query, memory_context) @@ -265,11 +260,8 @@ class TestBuildEnrichedQuery: """Test multiple context fields are appended.""" query = "What time and weather today?" memory_context = { - "profile": { - "location": "Amsterdam", - "timezone": "Europe/Amsterdam" - }, - "preferences": {"temperature_unit": "celsius"} + "profile": {"location": "Amsterdam", "timezone": "Europe/Amsterdam"}, + "preferences": {"temperature_unit": "celsius"}, } result = _build_enriched_query(query, memory_context) @@ -281,9 +273,7 @@ class TestBuildEnrichedQuery: def test_no_enrichment_for_unrelated_query(self): """Test no enrichment for queries that don't need context.""" query = "Tell me a joke" - memory_context = { - "profile": {"location": "Amsterdam", "timezone": "Europe/Amsterdam"} - } + memory_context = {"profile": {"location": "Amsterdam", "timezone": "Europe/Amsterdam"}} result = _build_enriched_query(query, memory_context) @@ -306,7 +296,10 @@ class TestExtractCapabilities: ("The user wants a description of the algorithm.", "script -> housekeeper"), ("I should discover what the answer is.", "cover -> housekeeper"), ("That sounds fantastic, let me compute it.", "fan -> housekeeper"), - ("I acknowledge the request to add two numbers.", "knowledge/know -> librarian, biographer"), + ( + "I acknowledge the request to add two numbers.", + "knowledge/know -> librarian, biographer", + ), ("The user asks about the economy myth.", "my -> biographer"), ("Convert 98.6 Fahrenheit to Celsius.", "temperature is a housekeeper domain"), ] diff --git a/tests/agents/test_delegation.py b/tests/agents/test_delegation.py index e0e6d4d..526a8a5 100644 --- a/tests/agents/test_delegation.py +++ b/tests/agents/test_delegation.py @@ -4,6 +4,7 @@ Tests for delegation infrastructure. Tests the DelegationTask dataclass and delegation wrapper functions that implement the agent-as-tool pattern. """ + from unittest.mock import AsyncMock, patch import pytest @@ -41,9 +42,7 @@ class TestBuildDelegationContext: assert "assistant: Docker is a container runtime." in context def test_only_last_max_turns_kept(self): - history = [ - {"role": "user", "content": f"message {i}"} for i in range(10) - ] + history = [{"role": "user", "content": f"message {i}"} for i in range(10)] context = build_delegation_context(history, max_turns=6) @@ -60,9 +59,7 @@ class TestBuildDelegationContext: assert "x" * 501 not in context def test_structured_content_parts_tolerated(self): - history = [ - {"role": "user", "content": [{"type": "text", "text": "hello there"}]} - ] + history = [{"role": "user", "content": [{"type": "text", "text": "hello there"}]}] context = build_delegation_context(history) @@ -343,8 +340,12 @@ class TestHouseholdThinkMessages: for action_type, messages in action_types.items(): for phase, msg in messages.items(): # Messages should NOT have wrappers - they go to reasoning_content field - assert "" not in msg, f"{expert}/{action_type}/{phase} should not have wrapper" - assert "" not in msg, f"{expert}/{action_type}/{phase} should not have wrapper" + assert ( + "" not in msg + ), f"{expert}/{action_type}/{phase} should not have wrapper" + assert ( + "" not in msg + ), f"{expert}/{action_type}/{phase} should not have wrapper" # Messages should be non-empty strings assert isinstance(msg, str) and len(msg) > 0, f"{expert}/{action_type}/{phase}" @@ -356,7 +357,9 @@ class TestDetectActionType: def test_librarian_search_is_retrieve(self): """Test librarian search tasks are RETRIEVE.""" assert _detect_action_type("librarian", "search for Docker info") == ActionType.RETRIEVE - assert _detect_action_type("librarian", "find information about CI/CD") == ActionType.RETRIEVE + assert ( + _detect_action_type("librarian", "find information about CI/CD") == ActionType.RETRIEVE + ) assert _detect_action_type("librarian", "look up Kubernetes docs") == ActionType.RETRIEVE def test_librarian_web_search_is_research(self): @@ -378,14 +381,25 @@ class TestDetectActionType: def test_biographer_record_is_record(self): """Test biographer record tasks are RECORD.""" - assert _detect_action_type("biographer", "remember that I work at Acme") == ActionType.RECORD + assert ( + _detect_action_type("biographer", "remember that I work at Acme") == ActionType.RECORD + ) assert _detect_action_type("biographer", "note that my car is a Tesla") == ActionType.RECORD - assert _detect_action_type("biographer", "save my preference for dark mode") == ActionType.RECORD + assert ( + _detect_action_type("biographer", "save my preference for dark mode") + == ActionType.RECORD + ) def test_housekeeper_status_is_retrieve(self): """Test housekeeper status tasks are RETRIEVE.""" - assert _detect_action_type("housekeeper", "what devices are in the bedroom?") == ActionType.RETRIEVE - assert _detect_action_type("housekeeper", "is the living room light on?") == ActionType.RETRIEVE + assert ( + _detect_action_type("housekeeper", "what devices are in the bedroom?") + == ActionType.RETRIEVE + ) + assert ( + _detect_action_type("housekeeper", "is the living room light on?") + == ActionType.RETRIEVE + ) def test_housekeeper_control_is_control(self): """Test housekeeper control tasks are CONTROL.""" diff --git a/tests/agents/test_lorem_tester.py b/tests/agents/test_lorem_tester.py index 6832f11..261936a 100644 --- a/tests/agents/test_lorem_tester.py +++ b/tests/agents/test_lorem_tester.py @@ -6,9 +6,9 @@ import pytest from src.agents.lorem_tester import LoremTesterAgent from src.core.exceptions import ( - RateLimitError, - ContextLengthError, APIError, + ContextLengthError, + RateLimitError, ) diff --git a/tests/agents/test_orchestration.py b/tests/agents/test_orchestration.py index d5c6b7b..918be0c 100644 --- a/tests/agents/test_orchestration.py +++ b/tests/agents/test_orchestration.py @@ -5,23 +5,25 @@ Tests the multi-expert coordination infrastructure including delegation parsing, think updates, result handling, and multi-expert sequential/parallel execution. """ -import pytest + from unittest.mock import AsyncMock, patch +import pytest + +from src.agents.delegation import DelegationResult, DelegationTask from src.agents.orchestration import ( - OrchestrationContext, - parse_delegation_from_steward_note, - execute_delegation, - orchestrate_with_think_updates, - extract_delegation_context, ExecutionMode, MultiExpertResult, - execute_sequential, - execute_parallel, - orchestrate_multi_expert, + OrchestrationContext, _get_display_name, + execute_delegation, + execute_parallel, + execute_sequential, + extract_delegation_context, + orchestrate_multi_expert, + orchestrate_with_think_updates, + parse_delegation_from_steward_note, ) -from src.agents.delegation import DelegationTask, DelegationResult @pytest.mark.unit @@ -364,6 +366,7 @@ class TestOrchestrationContext: # Multi-Expert Coordination Tests # ============================================================================ + @pytest.mark.unit class TestMultiExpertResult: """Tests for MultiExpertResult aggregation.""" @@ -414,18 +417,22 @@ class TestMultiExpertResult: """Test aggregating outputs from multiple experts.""" result = MultiExpertResult() - result.add_result(DelegationResult( - expert_name="librarian", - task="search docs", - success=True, - output="Found Docker docs", - )) - result.add_result(DelegationResult( - expert_name="memory", - task="get preferences", - success=True, - output="User prefers dark mode", - )) + result.add_result( + DelegationResult( + expert_name="librarian", + task="search docs", + success=True, + output="Found Docker docs", + ) + ) + result.add_result( + DelegationResult( + expert_name="memory", + task="get preferences", + success=True, + output="User prefers dark mode", + ) + ) combined = result.aggregate_outputs() @@ -438,19 +445,23 @@ class TestMultiExpertResult: """Test that failed results are excluded from aggregate.""" result = MultiExpertResult() - result.add_result(DelegationResult( - expert_name="librarian", - task="search", - success=True, - output="Success output", - )) - result.add_result(DelegationResult( - expert_name="memory", - task="get", - success=False, - output="", - error="Failed", - )) + result.add_result( + DelegationResult( + expert_name="librarian", + task="search", + success=True, + output="Success output", + ) + ) + result.add_result( + DelegationResult( + expert_name="memory", + task="get", + success=False, + output="", + error="Failed", + ) + ) combined = result.aggregate_outputs() @@ -471,7 +482,9 @@ class TestExecuteSequential: ] mock_results = [ - DelegationResult(expert_name="librarian", task="task 1", success=True, output="Result 1"), + DelegationResult( + expert_name="librarian", task="task 1", success=True, output="Result 1" + ), DelegationResult(expert_name="memory", task="task 2", success=True, output="Result 2"), ] @@ -496,7 +509,9 @@ class TestExecuteSequential: mock_results = [ DelegationResult(expert_name="librarian", task="task 1", success=True, output="OK"), - DelegationResult(expert_name="memory", task="task 2", success=False, output="", error="Failed"), + DelegationResult( + expert_name="memory", task="task 2", success=False, output="", error="Failed" + ), ] with patch( @@ -520,7 +535,9 @@ class TestExecuteSequential: ] mock_results = [ - DelegationResult(expert_name="librarian", task="task 1", success=False, output="", error="Error"), + DelegationResult( + expert_name="librarian", task="task 1", success=False, output="", error="Error" + ), ] with patch( @@ -548,7 +565,9 @@ class TestExecuteParallel: ] mock_results = [ - DelegationResult(expert_name="librarian", task="task 1", success=True, output="Result 1"), + DelegationResult( + expert_name="librarian", task="task 1", success=True, output="Result 1" + ), DelegationResult(expert_name="memory", task="task 2", success=True, output="Result 2"), ] @@ -572,7 +591,9 @@ class TestExecuteParallel: mock_results = [ DelegationResult(expert_name="librarian", task="task 1", success=True, output="OK"), - DelegationResult(expert_name="memory", task="task 2", success=False, output="", error="Timeout"), + DelegationResult( + expert_name="memory", task="task 2", success=False, output="", error="Timeout" + ), ] with patch( @@ -629,7 +650,9 @@ class TestOrchestrateMultiExpert: ] mock_results = [ - DelegationResult(expert_name="librarian", task="task 1", success=True, output="Result 1"), + DelegationResult( + expert_name="librarian", task="task 1", success=True, output="Result 1" + ), DelegationResult(expert_name="memory", task="task 2", success=True, output="Result 2"), ] @@ -658,7 +681,9 @@ class TestOrchestrateMultiExpert: ] mock_results = [ - DelegationResult(expert_name="librarian", task="task 1", success=True, output="Result 1"), + DelegationResult( + expert_name="librarian", task="task 1", success=True, output="Result 1" + ), DelegationResult(expert_name="memory", task="task 2", success=True, output="Result 2"), ] diff --git a/tests/agents/test_protocol.py b/tests/agents/test_protocol.py index 6ada5b1..9efc6b6 100644 --- a/tests/agents/test_protocol.py +++ b/tests/agents/test_protocol.py @@ -1,6 +1,7 @@ """ Tests for the agent error protocol. """ + import pytest from src.agents.protocol import AgentError diff --git a/tests/agents/test_registry.py b/tests/agents/test_registry.py index cb44a1a..ec062ab 100644 --- a/tests/agents/test_registry.py +++ b/tests/agents/test_registry.py @@ -4,8 +4,8 @@ Tests for model registry. import pytest -from src.agents.registry import ModelRegistry from src.agents.lorem_tester import LoremTesterAgent +from src.agents.registry import ModelRegistry from src.agents.tatlock import TatlockAgent from src.core.exceptions import ModelNotFoundError @@ -63,10 +63,10 @@ async def test_tatlock_capabilities(): # Tatlock Phase 1 - basic streaming, reasoning, and permanent tools assert capabilities["streaming"] is True - assert capabilities["reasoning"] is True # Basic reasoning summaries - assert capabilities["tools"] is True # Permanent tools: calculator, date/time, search - assert capabilities["vision"] is False # Future - assert capabilities["audio"] is False # Future + assert capabilities["reasoning"] is True # Basic reasoning summaries + assert capabilities["tools"] is True # Permanent tools: calculator, date/time, search + assert capabilities["vision"] is False # Future + assert capabilities["audio"] is False # Future @pytest.mark.unit diff --git a/tests/agents/test_tatlock_agent.py b/tests/agents/test_tatlock_agent.py index d235c28..12092db 100644 --- a/tests/agents/test_tatlock_agent.py +++ b/tests/agents/test_tatlock_agent.py @@ -6,9 +6,7 @@ These tests verify: 2. Tool calls are logged to reasoning output (users see what tools are doing) """ -import json import pytest -from unittest.mock import patch, AsyncMock from httpx import AsyncClient @@ -28,14 +26,10 @@ async def test_tatlock_conversation_history_memory(async_client: AsyncClient): "messages": [ {"role": "user", "content": "My name is Alice and I love Python programming."} ], - "stream": False + "stream": False, } - response_1 = await async_client.post( - "/v1/chat/completions", - json=request_data_1, - timeout=120.0 - ) + response_1 = await async_client.post("/v1/chat/completions", json=request_data_1, timeout=120.0) assert response_1.status_code == 200 data_1 = response_1.json() @@ -48,16 +42,15 @@ async def test_tatlock_conversation_history_memory(async_client: AsyncClient): "messages": [ {"role": "user", "content": "My name is Alice and I love Python programming."}, {"role": "assistant", "content": first_response}, - {"role": "user", "content": "What did I say my name was? And what programming language did I mention?"} + { + "role": "user", + "content": "What did I say my name was? And what programming language did I mention?", + }, ], - "stream": False + "stream": False, } - response_2 = await async_client.post( - "/v1/chat/completions", - json=request_data_2, - timeout=120.0 - ) + response_2 = await async_client.post("/v1/chat/completions", json=request_data_2, timeout=120.0) assert response_2.status_code == 200 data_2 = response_2.json() @@ -68,7 +61,9 @@ async def test_tatlock_conversation_history_memory(async_client: AsyncClient): has_python = "python" in second_response if not has_alice or not has_python: - pytest.xfail(f"LLM did not remember context (non-deterministic): alice={has_alice}, python={has_python}, response: {second_response[:200]}") + pytest.xfail( + f"LLM did not remember context (non-deterministic): alice={has_alice}, python={has_python}, response: {second_response[:200]}" + ) @pytest.mark.integration @@ -86,48 +81,37 @@ async def test_tatlock_multi_turn_context(async_client: AsyncClient): # Turn 1: Set up a topic conversation.append({"role": "user", "content": "Let's talk about the number 42."}) - request_1 = { - "model": "Tatlock", - "messages": conversation.copy(), - "stream": False - } + request_1 = {"model": "Tatlock", "messages": conversation.copy(), "stream": False} - response_1 = await async_client.post( - "/v1/chat/completions", - json=request_1, - timeout=120.0 - ) + response_1 = await async_client.post("/v1/chat/completions", json=request_1, timeout=120.0) assert response_1.status_code == 200 data_1 = response_1.json() - conversation.append({ - "role": "assistant", - "content": data_1["choices"][0]["message"]["content"] - }) + conversation.append( + {"role": "assistant", "content": data_1["choices"][0]["message"]["content"]} + ) # Turn 2: Reference "it" (should refer to 42) conversation.append({"role": "user", "content": "What number did I just mention?"}) - request_2 = { - "model": "Tatlock", - "messages": conversation.copy(), - "stream": False - } + request_2 = {"model": "Tatlock", "messages": conversation.copy(), "stream": False} - response_2 = await async_client.post( - "/v1/chat/completions", - json=request_2, - timeout=120.0 - ) + response_2 = await async_client.post("/v1/chat/completions", json=request_2, timeout=120.0) assert response_2.status_code == 200 data_2 = response_2.json() final_response = data_2["choices"][0]["message"]["content"] # Should reference 42 (check both as digit and word) - has_42 = "42" in final_response or "forty-two" in final_response.lower() or "forty two" in final_response.lower() + has_42 = ( + "42" in final_response + or "forty-two" in final_response.lower() + or "forty two" in final_response.lower() + ) if not has_42: - pytest.xfail(f"LLM did not mention 42 in response (non-deterministic): {final_response[:200]}") + pytest.xfail( + f"LLM did not mention 42 in response (non-deterministic): {final_response[:200]}" + ) @pytest.mark.integration @@ -142,16 +126,15 @@ async def test_tatlock_tool_call_logging_search(async_client: AsyncClient): request_data = { "model": "Tatlock", "messages": [ - {"role": "user", "content": "Search for current information about Python 3.13 release date"} + { + "role": "user", + "content": "Search for current information about Python 3.13 release date", + } ], - "stream": False + "stream": False, } - response = await async_client.post( - "/v1/chat/completions", - json=request_data, - timeout=120.0 - ) + response = await async_client.post("/v1/chat/completions", json=request_data, timeout=120.0) assert response.status_code == 200 data = response.json() @@ -167,8 +150,9 @@ async def test_tatlock_tool_call_logging_search(async_client: AsyncClient): # If search was used, should show the 🔍 emoji if "🔍" in full_response: - assert "search" in full_response.lower() or "python" in full_response.lower(), \ - "Search query should be visible in the response" + assert ( + "search" in full_response.lower() or "python" in full_response.lower() + ), "Search query should be visible in the response" @pytest.mark.integration @@ -183,45 +167,38 @@ async def test_tatlock_tool_call_logging_calculator(async_client: AsyncClient): """ request_data = { "model": "Tatlock", - "messages": [ - {"role": "user", "content": "What is the square root of 144 plus 25?"} - ], - "stream": False + "messages": [{"role": "user", "content": "What is the square root of 144 plus 25?"}], + "stream": False, } - response = await async_client.post( - "/v1/chat/completions", - json=request_data, - timeout=120.0 - ) + response = await async_client.post("/v1/chat/completions", json=request_data, timeout=120.0) assert response.status_code == 200 data = response.json() full_response = data["choices"][0]["message"]["content"] # Should have reasoning in tags (from Steward analysis) - assert "" in full_response, \ - f"Should have reasoning output in tags. Got: {full_response}" + assert ( + "" in full_response + ), f"Should have reasoning output in tags. Got: {full_response}" # Should reference the calculation in some form has_calculation_reference = ( - "144" in full_response or - "sqrt" in full_response.lower() or - "square root" in full_response.lower() + "144" in full_response + or "sqrt" in full_response.lower() + or "square root" in full_response.lower() ) - assert has_calculation_reference, \ - f"Should reference the calculation. Got: {full_response}" + assert has_calculation_reference, f"Should reference the calculation. Got: {full_response}" # Should have the correct answer (37) - assert "37" in full_response, \ - f"Should contain the answer 37. Got: {full_response}" + assert "37" in full_response, f"Should contain the answer 37. Got: {full_response}" # Tool emoji is optional - depends on whether tool was used directly # or computation was delegated to capability if "🧮" in full_response: - print(f"\nCalculator tool was used directly") + print("\nCalculator tool was used directly") else: - print(f"\nCalculation handled via tatlock_core capability") + print("\nCalculation handled via tatlock_core capability") print(f"\nCalculator response: {full_response}") @@ -234,17 +211,11 @@ async def test_tatlock_tool_call_logging_datetime(async_client: AsyncClient): """ request_data = { "model": "Tatlock", - "messages": [ - {"role": "user", "content": "What was the date exactly 2 weeks ago?"} - ], - "stream": False + "messages": [{"role": "user", "content": "What was the date exactly 2 weeks ago?"}], + "stream": False, } - response = await async_client.post( - "/v1/chat/completions", - json=request_data, - timeout=120.0 - ) + response = await async_client.post("/v1/chat/completions", json=request_data, timeout=120.0) assert response.status_code == 200 data = response.json() @@ -257,19 +228,36 @@ async def test_tatlock_tool_call_logging_datetime(async_client: AsyncClient): used_date_tool = "🕐" in full_response # Should mention the calculation or the timeframe - assert "2 weeks ago" in full_response.lower() or "weeks" in full_response.lower(), \ - f"Should reference the requested timeframe. Got: {full_response}" + assert ( + "2 weeks ago" in full_response.lower() or "weeks" in full_response.lower() + ), f"Should reference the requested timeframe. Got: {full_response}" # Should provide a specific date (either YYYY-MM-DD format or natural language like "November 23") import re - has_iso_date = bool(re.search(r'\d{4}-\d{2}-\d{2}', full_response)) - has_month_mention = any(month in full_response.lower() for month in - ['january', 'february', 'march', 'april', 'may', 'june', - 'july', 'august', 'september', 'october', 'november', 'december']) - has_date_number = bool(re.search(r'\b\d{1,2}(st|nd|rd|th)?\b', full_response.lower())) - assert has_iso_date or has_month_mention or has_date_number, \ - f"Should contain a specific date. Got: {full_response}" + has_iso_date = bool(re.search(r"\d{4}-\d{2}-\d{2}", full_response)) + has_month_mention = any( + month in full_response.lower() + for month in [ + "january", + "february", + "march", + "april", + "may", + "june", + "july", + "august", + "september", + "october", + "november", + "december", + ] + ) + has_date_number = bool(re.search(r"\b\d{1,2}(st|nd|rd|th)?\b", full_response.lower())) + + assert ( + has_iso_date or has_month_mention or has_date_number + ), f"Should contain a specific date. Got: {full_response}" print(f"\nDate/time response (tool used: {used_date_tool}): {full_response}") @@ -284,17 +272,11 @@ async def test_tatlock_no_tool_calls_no_logging(async_client: AsyncClient): """ request_data = { "model": "Tatlock", - "messages": [ - {"role": "user", "content": "Just say hello to me."} - ], - "stream": False + "messages": [{"role": "user", "content": "Just say hello to me."}], + "stream": False, } - response = await async_client.post( - "/v1/chat/completions", - json=request_data, - timeout=120.0 - ) + response = await async_client.post("/v1/chat/completions", json=request_data, timeout=120.0) assert response.status_code == 200 data = response.json() @@ -327,17 +309,9 @@ async def test_tatlock_conversation_history_with_tools(async_client: AsyncClient # Turn 1: Do a calculation conversation.append({"role": "user", "content": "Calculate 15 times 7 for me."}) - request_1 = { - "model": "Tatlock", - "messages": conversation.copy(), - "stream": False - } + request_1 = {"model": "Tatlock", "messages": conversation.copy(), "stream": False} - response_1 = await async_client.post( - "/v1/chat/completions", - json=request_1, - timeout=120.0 - ) + response_1 = await async_client.post("/v1/chat/completions", json=request_1, timeout=120.0) assert response_1.status_code == 200 data_1 = response_1.json() @@ -353,17 +327,9 @@ async def test_tatlock_conversation_history_with_tools(async_client: AsyncClient # Turn 2: Ask about previous calculation conversation.append({"role": "user", "content": "What calculation did I just ask you to do?"}) - request_2 = { - "model": "Tatlock", - "messages": conversation.copy(), - "stream": False - } + request_2 = {"model": "Tatlock", "messages": conversation.copy(), "stream": False} - response_2 = await async_client.post( - "/v1/chat/completions", - json=request_2, - timeout=120.0 - ) + response_2 = await async_client.post("/v1/chat/completions", json=request_2, timeout=120.0) assert response_2.status_code == 200 data_2 = response_2.json() @@ -371,13 +337,15 @@ async def test_tatlock_conversation_history_with_tools(async_client: AsyncClient # Should remember the calculation (either as digits or words) has_calculation = ( - ("15" in second_response and "7" in second_response) or # As digits - ("fifteen" in second_response and "seven" in second_response) or # As words - "105" in second_response or # As answer - "multipl" in second_response # Mentions multiplication + ("15" in second_response and "7" in second_response) # As digits + or ("fifteen" in second_response and "seven" in second_response) # As words + or "105" in second_response # As answer + or "multipl" in second_response # Mentions multiplication ) if not has_calculation: - pytest.xfail(f"LLM did not remember calculation (non-deterministic): {second_response[:200]}") + pytest.xfail( + f"LLM did not remember calculation (non-deterministic): {second_response[:200]}" + ) @pytest.mark.integration @@ -404,19 +372,13 @@ async def test_tatlock_ollama_fallback(async_client: AsyncClient): request_data = { "model": "Tatlock", - "messages": [ - {"role": "user", "content": "Say hello to me."} - ], - "stream": False + "messages": [{"role": "user", "content": "Say hello to me."}], + "stream": False, } # 300s: this test forbids the Claude rescue, and the full local # Steward -> orchestrate -> synthesize flow on gemma4 exceeds 120s - response = await async_client.post( - "/v1/chat/completions", - json=request_data, - timeout=300.0 - ) + response = await async_client.post("/v1/chat/completions", json=request_data, timeout=300.0) assert response.status_code == 200 data = response.json() diff --git a/tests/agents/test_tools.py b/tests/agents/test_tools.py index dcba341..cb8de99 100644 --- a/tests/agents/test_tools.py +++ b/tests/agents/test_tools.py @@ -5,21 +5,20 @@ Note: Web search has been moved to The Librarian agent. See tests/agents/librarian/test_tools.py for search tests. """ -import pytest from datetime import datetime from src.agents.tools import ( calculate, - get_current_datetime, calculate_time_offset, + get_current_datetime, time_difference, ) - # ============================================================================ # Calculator Tests # ============================================================================ + class TestCalculator: """Tests for the calculator tool.""" @@ -81,6 +80,7 @@ class TestCalculator: # Date/Time Tests # ============================================================================ + class TestDateTime: """Tests for date/time toolkit.""" diff --git a/tests/anthropic/test_model_selector.py b/tests/anthropic/test_model_selector.py index b0f38cd..0a3086f 100644 --- a/tests/anthropic/test_model_selector.py +++ b/tests/anthropic/test_model_selector.py @@ -4,6 +4,7 @@ Unit tests for backend selection (Ollama primary, Claude fallback). These tests set the cached health-check globals directly so they are deterministic regardless of which services are reachable. """ + import pytest from src.anthropic import model_selector diff --git a/tests/chat/test_router.py b/tests/chat/test_router.py index 4985168..b8cb96c 100644 --- a/tests/chat/test_router.py +++ b/tests/chat/test_router.py @@ -1,6 +1,7 @@ """ Tests for chat completions router. """ + import json import pytest @@ -17,24 +18,24 @@ def test_chat_completion_non_streaming( ) -> None: """Test non-streaming chat completion.""" response = client.post("/v1/chat/completions", json=mock_chat_request) - + assert response.status_code == 200 data = response.json() - + # Verify response structure assert "id" in data assert data["object"] == constants.CHAT_COMPLETION_OBJECT assert "created" in data assert data["model"] == mock_chat_request["model"] assert len(data["choices"]) == 1 - + # Verify choice structure choice = data["choices"][0] assert choice["index"] == 0 assert choice["message"]["role"] == constants.ROLE_ASSISTANT assert choice["message"]["content"] # Should have content assert choice["finish_reason"] == constants.FINISH_REASON_STOP - + # Verify usage assert "usage" in data assert data["usage"]["prompt_tokens"] > 0 @@ -47,9 +48,9 @@ def test_chat_completion_validation_error(client: TestClient) -> None: """Test chat completion with invalid request.""" # Missing required field 'messages' invalid_request = {"model": "Tatlock"} - + response = client.post("/v1/chat/completions", json=invalid_request) - + assert response.status_code == 422 data = response.json() assert "error" in data @@ -108,7 +109,7 @@ async def test_chat_completion_streaming( # Execute with overall timeout try: chunks = await asyncio.wait_for(read_stream_with_timeout(), timeout=20.0) - except asyncio.TimeoutError: + except TimeoutError: pytest.fail("Streaming test timed out after 20 seconds") # Verify we got chunks @@ -133,7 +134,7 @@ def test_chat_completion_temperature_validation( invalid_request = {**mock_chat_request, "temperature": 3.0} response = client.post("/v1/chat/completions", json=invalid_request) assert response.status_code == 422 - + # Valid temperature valid_request = {**mock_chat_request, "temperature": 0.5} response = client.post("/v1/chat/completions", json=valid_request) @@ -153,7 +154,7 @@ def test_chat_completion_message_roles( {"role": "user", "content": "Hello!"}, ], } - + response = client.post("/v1/chat/completions", json=request_with_system) assert response.status_code == 200 @@ -168,6 +169,6 @@ def test_chat_completion_invalid_role( **mock_chat_request, "messages": [{"role": "invalid_role", "content": "test"}], } - + response = client.post("/v1/chat/completions", json=invalid_request) assert response.status_code == 422 diff --git a/tests/chat/test_streaming_wrapper.py b/tests/chat/test_streaming_wrapper.py index 964b440..29fc6e6 100644 --- a/tests/chat/test_streaming_wrapper.py +++ b/tests/chat/test_streaming_wrapper.py @@ -7,7 +7,9 @@ Tests that the wrapper correctly: - Streams reasoning via reasoning_content field (DeepSeek R1 format) - Streams both reasoning and content """ + import json + import pytest from httpx import AsyncClient @@ -20,10 +22,8 @@ async def test_streaming_wrapper_enables_reasoning(async_client: AsyncClient): """Test that streaming wrapper automatically enables reasoning via reasoning_content.""" request_data = { "model": "lorem-tester", - "messages": [ - {"role": "user", "content": "Test message"} - ], - "stream": True + "messages": [{"role": "user", "content": "Test message"}], + "stream": True, } chunks_received = [] @@ -74,10 +74,8 @@ async def test_streaming_wrapper_reasoning_before_content(async_client: AsyncCli """Test that reasoning_content comes before regular content.""" request_data = { "model": "lorem-tester", - "messages": [ - {"role": "user", "content": "Explain something"} - ], - "stream": True + "messages": [{"role": "user", "content": "Explain something"}], + "stream": True, } chunk_types = [] # Track order: 'reasoning' or 'content' @@ -127,11 +125,9 @@ async def test_streaming_wrapper_proper_chunk_structure(async_client: AsyncClien """Test that streaming chunks have proper structure.""" request_data = { "model": "lorem-tester", - "messages": [ - {"role": "user", "content": "Hello"} - ], + "messages": [{"role": "user", "content": "Hello"}], "temperature": 0.8, - "stream": True + "stream": True, } first_chunk = None @@ -201,9 +197,9 @@ async def test_streaming_wrapper_with_system_message(async_client: AsyncClient): "model": "lorem-tester", "messages": [ {"role": "system", "content": "You are a helpful assistant."}, - {"role": "user", "content": "Hello"} + {"role": "user", "content": "Hello"}, ], - "stream": True + "stream": True, } chunks_received = [] @@ -243,10 +239,8 @@ async def test_streaming_wrapper_pipeline_prefix(async_client: AsyncClient): """Test streaming with pipeline prefix in model name.""" request_data = { "model": "some_pipeline.lorem-tester", - "messages": [ - {"role": "user", "content": "Test"} - ], - "stream": True + "messages": [{"role": "user", "content": "Test"}], + "stream": True, } chunks_received = [] @@ -285,17 +279,11 @@ async def test_streaming_wrapper_non_streaming_fallback(async_client: AsyncClien """Test that non-streaming request works through wrapper.""" request_data = { "model": "lorem-tester", - "messages": [ - {"role": "user", "content": "Hello"} - ], - "stream": False # Non-streaming + "messages": [{"role": "user", "content": "Hello"}], + "stream": False, # Non-streaming } - response = await async_client.post( - "/v1/chat/completions", - json=request_data, - timeout=20.0 - ) + response = await async_client.post("/v1/chat/completions", json=request_data, timeout=20.0) assert response.status_code == 200 data = response.json() diff --git a/tests/conftest.py b/tests/conftest.py index 68ce14a..59b334e 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -2,11 +2,12 @@ Shared test fixtures for all tests. Following FastAPI testing best practices. """ + import asyncio import pytest from fastapi.testclient import TestClient -from httpx import AsyncClient, ASGITransport +from httpx import ASGITransport, AsyncClient from src.main import app @@ -28,10 +29,7 @@ def _tenant_guard(): from src.core.multi_tenancy import get_memory_collection_name effective = get_default_user() - if ( - effective == PRODUCTION_TENANT - or config.effective_default_user == PRODUCTION_TENANT - ): + if effective == PRODUCTION_TENANT or config.effective_default_user == PRODUCTION_TENANT: pytest.exit( f"TENANT GUARD: refusing to run the test suite - the effective " f"tenant resolves to the production tenant '{PRODUCTION_TENANT}' " @@ -59,6 +57,7 @@ def _initialize_app(_tenant_guard): production tenant before any initialization happens. """ from src.core.startup import initialize_application + asyncio.run(initialize_application()) @@ -79,10 +78,7 @@ async def async_client() -> AsyncClient: Use for testing async endpoints and streaming. """ - async with AsyncClient( - transport=ASGITransport(app=app), - base_url="http://test" - ) as client: + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: yield client @@ -91,9 +87,7 @@ def mock_chat_request() -> dict: """Standard chat completion request fixture.""" return { "model": "lorem-tester", - "messages": [ - {"role": "user", "content": "Hello, world!"} - ], + "messages": [{"role": "user", "content": "Hello, world!"}], "temperature": 0.7, "stream": False, } diff --git a/tests/contracts/test_service_contracts.py b/tests/contracts/test_service_contracts.py index 4cadd8e..1c5db2d 100644 --- a/tests/contracts/test_service_contracts.py +++ b/tests/contracts/test_service_contracts.py @@ -11,6 +11,7 @@ Semantics: Run with: make test-contracts """ + import json import httpx @@ -66,9 +67,9 @@ class TestOllamaContract: assert response.status_code == 200 names = [m["name"] for m in response.json()["models"]] model = config.OLLAMA_DEFAULT_MODEL - assert model in names or f"{model}:latest" in names, ( - f"{model} not pulled; available: {names}" - ) + assert ( + model in names or f"{model}:latest" in names + ), f"{model} not pulled; available: {names}" async def test_generate_returns_plain_text(self): # Mirrors StewardAgent._call_ollama() @@ -94,9 +95,7 @@ class TestOllamaContract: "ollama", { "model": config.OLLAMA_DEFAULT_MODEL, - "messages": [ - {"role": "user", "content": "What is 6 * 7? Use the calculator."} - ], + "messages": [{"role": "user", "content": "What is 6 * 7? Use the calculator."}], "tools": [CALCULATOR_TOOL], "tool_choice": "required", "stream": False, diff --git a/tests/core/test_household_registry.py b/tests/core/test_household_registry.py index 4783fc8..b42e58d 100644 --- a/tests/core/test_household_registry.py +++ b/tests/core/test_household_registry.py @@ -3,6 +3,7 @@ Tests for household registry. Tests capability registration, toolset scoping, and coordination features. """ + import pytest from pydantic_ai.tools import Tool @@ -10,7 +11,6 @@ from src.core.household_registry import ( HouseholdCapability, HouseholdMember, HouseholdRegistry, - household_registry, ) @@ -38,6 +38,7 @@ def sample_capability(): @pytest.fixture def sample_tools(): """Sample tool definitions.""" + def test_function_1(x: int) -> int: """Test function 1.""" return x * 2 @@ -96,6 +97,7 @@ class TestHouseholdMember: def test_member_with_agent(self, sample_capability, sample_tools): """Test member can include an agent.""" from unittest.mock import Mock + mock_agent = Mock() member = HouseholdMember( @@ -322,7 +324,9 @@ class TestGetDelegationTools: assert callable(tools[0]) assert tools[0].__name__ == "delegate_to_librarian" - def test_delegation_tools_returns_raw_tools_for_member_without_agent(self, registry, sample_capability, sample_tools): + def test_delegation_tools_returns_raw_tools_for_member_without_agent( + self, registry, sample_capability, sample_tools + ): """Test delegation tools returns raw tools when member has no agent.""" registry.register("test_tools", sample_capability, sample_tools) @@ -373,8 +377,8 @@ class TestGetDelegationTools: assert tools[0].__name__ == "delegate_to_librarian" # Rest should be raw tools - assert hasattr(tools[1], 'name') - assert hasattr(tools[2], 'name') + assert hasattr(tools[1], "name") + assert hasattr(tools[2], "name") def test_delegation_tools_nonexistent_member(self, registry): """Test delegation tools handles non-existent member gracefully.""" diff --git a/tests/core/test_logging_config.py b/tests/core/test_logging_config.py index d7db179..2483360 100644 --- a/tests/core/test_logging_config.py +++ b/tests/core/test_logging_config.py @@ -3,12 +3,11 @@ Tests for structured logging configuration. Tests logging setup, context management, and FastAPI integration. """ + import logging -from io import StringIO from unittest.mock import patch import pytest -import structlog from src.core.logging_config import ( add_log_level, @@ -50,10 +49,10 @@ class TestGetLogger: """Test get_logger returns structlog BoundLogger.""" logger = get_logger("test") # Logger should have standard logging methods - assert hasattr(logger, 'info') - assert hasattr(logger, 'debug') - assert hasattr(logger, 'warning') - assert hasattr(logger, 'error') + assert hasattr(logger, "info") + assert hasattr(logger, "debug") + assert hasattr(logger, "warning") + assert hasattr(logger, "error") def test_get_logger_with_module_name(self): """Test logger with module name.""" diff --git a/tests/core/test_memory_service.py b/tests/core/test_memory_service.py index 974aa6a..da07204 100644 --- a/tests/core/test_memory_service.py +++ b/tests/core/test_memory_service.py @@ -2,13 +2,14 @@ Tests for the memory service (direct access layer). """ +from unittest.mock import AsyncMock, patch + import pytest -from unittest.mock import MagicMock, patch, AsyncMock from src.core.memory_service import ( + MemoryRecord, MemoryService, MemoryType, - MemoryRecord, memory_service, ) diff --git a/tests/core/test_router.py b/tests/core/test_router.py index 4e9b787..481d41e 100644 --- a/tests/core/test_router.py +++ b/tests/core/test_router.py @@ -1,6 +1,7 @@ """ Tests for core router (health check and root endpoints). """ + import pytest from fastapi.testclient import TestClient @@ -11,7 +12,7 @@ from src.core.config import config def test_health_check(client: TestClient) -> None: """Test health check endpoint returns healthy status.""" response = client.get("/health") - + assert response.status_code == 200 data = response.json() assert data["status"] == "healthy" @@ -22,7 +23,7 @@ def test_health_check(client: TestClient) -> None: def test_root_endpoint(client: TestClient) -> None: """Test root endpoint returns API information.""" response = client.get("/") - + assert response.status_code == 200 data = response.json() assert data["name"] == config.APP_NAME @@ -42,7 +43,7 @@ def test_openapi_schema(client: TestClient) -> None: """Test that OpenAPI schema is accessible.""" response = client.get("/openapi.json") assert response.status_code == 200 - + schema = response.json() assert schema["info"]["title"] == config.APP_NAME assert schema["info"]["version"] == config.APP_VERSION diff --git a/tests/core/test_tenant_guard.py b/tests/core/test_tenant_guard.py index 7da9d3b..f730ce7 100644 --- a/tests/core/test_tenant_guard.py +++ b/tests/core/test_tenant_guard.py @@ -157,9 +157,7 @@ class TestRequestContextGuard: "jpmschweitzer!", ], ) - @pytest.mark.parametrize( - "environment", [Environment.DEVELOPMENT, Environment.TESTING] - ) + @pytest.mark.parametrize("environment", [Environment.DEVELOPMENT, Environment.TESTING]) def test_production_tenant_sanitization_variants_are_forced( self, monkeypatch, environment, variant ): @@ -175,9 +173,8 @@ class TestRequestContextGuard: with RequestContext(user=variant): effective = get_user() assert effective == TEST_TENANT - assert ( - get_memory_collection_name(effective) - != get_memory_collection_name(PRODUCTION_TENANT) + assert get_memory_collection_name(effective) != get_memory_collection_name( + PRODUCTION_TENANT ) def test_dev_non_colliding_user_is_not_forced(self, monkeypatch): @@ -228,10 +225,7 @@ class TestSuiteRunsUnderTestTenant: from src.core.context import get_default_user from src.core.multi_tenancy import get_memory_collection_name - assert ( - get_memory_collection_name(get_default_user()) - == f"memories_{TEST_TENANT}" - ) + assert get_memory_collection_name(get_default_user()) == f"memories_{TEST_TENANT}" def test_redis_session_namespace_is_test_tenant(self): from src.core.context import get_default_user diff --git a/tests/core/test_tool_tracking.py b/tests/core/test_tool_tracking.py index 947b5cc..a45030e 100644 --- a/tests/core/test_tool_tracking.py +++ b/tests/core/test_tool_tracking.py @@ -3,6 +3,7 @@ Tests for tool call tracking. Tests capability extraction and recommendation matching. """ + import pytest from src.core.tool_tracking import ToolCallTracker @@ -29,9 +30,7 @@ class TestToolCallTracker: @pytest.mark.asyncio async def test_track_call_recognizes_delegation_as_recommended(self): """Test that delegate_to_X is recognized when X is recommended.""" - tracker = ToolCallTracker( - recommended_capabilities=["librarian", "biographer"] - ) + tracker = ToolCallTracker(recommended_capabilities=["librarian", "biographer"]) await tracker.track_call("delegate_to_librarian", 1.0) @@ -42,9 +41,7 @@ class TestToolCallTracker: @pytest.mark.asyncio async def test_track_call_detects_not_recommended(self): """Test that unrecommended tools are flagged.""" - tracker = ToolCallTracker( - recommended_capabilities=["librarian"] - ) + tracker = ToolCallTracker(recommended_capabilities=["librarian"]) await tracker.track_call("delegate_to_housekeeper", 1.0) @@ -55,9 +52,7 @@ class TestToolCallTracker: def test_get_summary_with_delegation_tools(self): """Test summary correctly maps delegation tools to capabilities.""" - tracker = ToolCallTracker( - recommended_capabilities=["librarian", "biographer"] - ) + tracker = ToolCallTracker(recommended_capabilities=["librarian", "biographer"]) tracker.actual_calls = { "delegate_to_librarian": [1.0, 2.0], "delegate_to_housekeeper": [0.5], # Not recommended @@ -72,9 +67,7 @@ class TestToolCallTracker: @pytest.mark.asyncio async def test_finalize_with_delegation_tools(self): """Test finalize correctly identifies unused recommendations.""" - tracker = ToolCallTracker( - recommended_capabilities=["librarian", "biographer"] - ) + tracker = ToolCallTracker(recommended_capabilities=["librarian", "biographer"]) tracker.actual_calls = { "delegate_to_librarian": [1.0], } diff --git a/tests/e2e/test_api_endpoints.py b/tests/e2e/test_api_endpoints.py index 1f9cc72..c3b84b6 100644 --- a/tests/e2e/test_api_endpoints.py +++ b/tests/e2e/test_api_endpoints.py @@ -7,10 +7,12 @@ These tests hit the actual running server and test the full stack: - Tool execution - Response formatting """ + +from collections.abc import AsyncGenerator + +import httpx import pytest import pytest_asyncio -import httpx -from typing import AsyncGenerator # Test server base URL (assumes server is running on localhost:8777 via ./wakeup.sh) BASE_URL = "http://localhost:8777" @@ -34,10 +36,8 @@ class TestChatCompletionsE2E: "/v1/chat/completions", json={ "model": "Tatlock", - "messages": [ - {"role": "user", "content": "What is 144 divided by 12?"} - ], - } + "messages": [{"role": "user", "content": "What is 144 divided by 12?"}], + }, ) assert response.status_code == 200 @@ -68,7 +68,9 @@ class TestChatCompletionsE2E: assert "usage" in data assert data["usage"]["total_tokens"] > 0 - print(f"✓ Calculator test passed. Found '12' in response. Tool indicator: {has_calculator_indicator}") + print( + f"✓ Calculator test passed. Found '12' in response. Tool indicator: {has_calculator_indicator}" + ) @pytest.mark.asyncio async def test_web_search(self, client: httpx.AsyncClient): @@ -80,7 +82,7 @@ class TestChatCompletionsE2E: "messages": [ {"role": "user", "content": "Search for the current population of Tokyo"} ], - } + }, ) assert response.status_code == 200 @@ -114,10 +116,8 @@ class TestChatCompletionsE2E: "/v1/chat/completions", json={ "model": "Tatlock", - "messages": [ - {"role": "user", "content": "What is 15 times 4?"} - ], - } + "messages": [{"role": "user", "content": "What is 15 times 4?"}], + }, ) assert response1.status_code == 200 @@ -135,9 +135,9 @@ class TestChatCompletionsE2E: "messages": [ {"role": "user", "content": "What is 15 times 4?"}, {"role": "assistant", "content": message1}, - {"role": "user", "content": "Now add 20 to that result."} + {"role": "user", "content": "Now add 20 to that result."}, ], - } + }, ) assert response2.status_code == 200 @@ -152,7 +152,9 @@ class TestChatCompletionsE2E: has_calculation = "60" in message2 and "20" in message2 assert has_answer or has_calculation, f"Expected '80' or calculation in: {message2}" - print(f"✓ Multi-turn test passed. Answer found: {has_answer}, Calculation shown: {has_calculation}") + print( + f"✓ Multi-turn test passed. Answer found: {has_answer}, Calculation shown: {has_calculation}" + ) @pytest.mark.asyncio async def test_calculation_and_search(self, client: httpx.AsyncClient): @@ -161,13 +163,8 @@ class TestChatCompletionsE2E: "/v1/chat/completions", json={ "model": "Tatlock", - "messages": [ - { - "role": "user", - "content": "Calculate the square root of 256" - } - ], - } + "messages": [{"role": "user", "content": "Calculate the square root of 256"}], + }, ) assert response.status_code == 200 @@ -194,10 +191,8 @@ class TestChatCompletionsE2E: "/v1/chat/completions", json={ "model": "Tatlock", - "messages": [ - {"role": "user", "content": "Hello, how are you?"} - ], - } + "messages": [{"role": "user", "content": "Hello, how are you?"}], + }, ) assert response.status_code == 200 @@ -223,10 +218,8 @@ class TestChatCompletionsE2E: "/v1/chat/completions", json={ "model": "Tatlock", - "messages": [ - {"role": "user", "content": "What is today's date?"} - ], - } + "messages": [{"role": "user", "content": "What is today's date?"}], + }, ) assert response.status_code == 200 @@ -242,11 +235,16 @@ class TestChatCompletionsE2E: # Should contain some date/time information (flexible - varies in format) import re + has_date = ( - re.search(r'\d{4}', message) or # Year - re.search(r'\d{1,2}', message) or # Day/month number - re.search(r'(January|February|March|April|May|June|July|August|September|October|November|December)', message, re.IGNORECASE) or - "today" in message.lower() + re.search(r"\d{4}", message) # Year + or re.search(r"\d{1,2}", message) # Day/month number + or re.search( + r"(January|February|March|April|May|June|July|August|September|October|November|December)", + message, + re.IGNORECASE, + ) + or "today" in message.lower() ) assert has_date, f"Expected date/time information in: {message}" @@ -263,11 +261,9 @@ class TestResponsesAPIE2E: "/v1/responses", json={ "model": "Tatlock", - "input": [ - {"role": "user", "content": "Calculate 25 times 16"} - ], - "reasoning": {"effort": "medium", "summary": "auto"} - } + "input": [{"role": "user", "content": "Calculate 25 times 16"}], + "reasoning": {"effort": "medium", "summary": "auto"}, + }, ) assert response.status_code == 200 @@ -300,7 +296,7 @@ class TestResponsesAPIE2E: assert "usage" in data assert data["usage"]["total_tokens"] > 0 - print(f"✓ Responses API test passed. Found '400' with Steward reasoning.") + print("✓ Responses API test passed. Found '400' with Steward reasoning.") @pytest.mark.asyncio async def test_response_multi_turn(self, client: httpx.AsyncClient): @@ -312,10 +308,10 @@ class TestResponsesAPIE2E: "input": [ {"role": "user", "content": "What is 7 times 8?"}, {"role": "assistant", "content": "Certainly, sir. 7 times 8 equals 56."}, - {"role": "user", "content": "Double that number."} + {"role": "user", "content": "Double that number."}, ], - "reasoning": {"effort": "medium", "summary": "auto"} - } + "reasoning": {"effort": "medium", "summary": "auto"}, + }, ) assert response.status_code == 200 @@ -346,11 +342,9 @@ class TestStreamingE2E: "/v1/chat/completions", json={ "model": "Tatlock", - "messages": [ - {"role": "user", "content": "What is 9 times 7?"} - ], - "stream": True - } + "messages": [{"role": "user", "content": "What is 9 times 7?"}], + "stream": True, + }, ) as response: assert response.status_code == 200 @@ -362,6 +356,7 @@ class TestStreamingE2E: break import json + chunk = json.loads(data_str) chunks.append(chunk) @@ -373,8 +368,7 @@ class TestStreamingE2E: # Should have received Steward's reasoning (in tags) full_content = "".join( - chunk["choices"][0]["delta"].get("content", "") or "" - for chunk in chunks + chunk["choices"][0]["delta"].get("content", "") or "" for chunk in chunks ) assert "" in full_content assert "" in full_content @@ -393,10 +387,8 @@ class TestErrorHandling: "/v1/chat/completions", json={ "model": "nonexistent-model", - "messages": [ - {"role": "user", "content": "Hello"} - ], - } + "messages": [{"role": "user", "content": "Hello"}], + }, ) assert response.status_code == 404 @@ -411,7 +403,7 @@ class TestErrorHandling: json={ "model": "Tatlock", # Missing "messages" field - } + }, ) assert response.status_code == 422 @@ -425,11 +417,9 @@ class TestErrorHandling: "/v1/chat/completions", json={ "model": "Tatlock", - "messages": [ - {"role": "user", "content": "Hello"} - ], - "temperature": 5.0 # Max is 2.0 - } + "messages": [{"role": "user", "content": "Hello"}], + "temperature": 5.0, # Max is 2.0 + }, ) assert response.status_code == 422 @@ -447,11 +437,9 @@ class TestChatResponsesWrapper: "/v1/responses", json={ "model": "Tatlock", - "input": [ - {"role": "user", "content": "Calculate 13 times 9"} - ], - "reasoning": {"effort": "medium", "summary": "auto"} - } + "input": [{"role": "user", "content": "Calculate 13 times 9"}], + "reasoning": {"effort": "medium", "summary": "auto"}, + }, ) assert response.status_code == 200 @@ -500,10 +488,8 @@ class TestChatResponsesWrapper: "/v1/chat/completions", json={ "model": "Tatlock", - "messages": [ - {"role": "user", "content": "What is 5 plus 3?"} - ], - } + "messages": [{"role": "user", "content": "What is 5 plus 3?"}], + }, ) assert response.status_code == 200 @@ -539,11 +525,9 @@ class TestChatResponsesWrapper: "/v1/chat/completions", json={ "model": "Tatlock", - "messages": [ - {"role": "user", "content": "Count to 3"} - ], - "stream": True - } + "messages": [{"role": "user", "content": "Count to 3"}], + "stream": True, + }, ) as response: assert response.status_code == 200 @@ -555,6 +539,7 @@ class TestChatResponsesWrapper: break import json + chunk = json.loads(data_str) chunks.append(chunk) @@ -574,10 +559,7 @@ class TestChatResponsesWrapper: assert chunks[0]["choices"][0]["delta"]["role"] == "assistant" # Should have content chunks - has_content = any( - "content" in chunk["choices"][0]["delta"] - for chunk in chunks - ) + has_content = any("content" in chunk["choices"][0]["delta"] for chunk in chunks) assert has_content print(f"✓ Streaming format matches OpenAI spec ({len(chunks)} chunks)") @@ -593,11 +575,9 @@ class TestStewardIntegration: "/v1/responses", json={ "model": "Tatlock", - "input": [ - {"role": "user", "content": "Calculate 123 times 456"} - ], - "reasoning": {"effort": "medium", "summary": "auto"} - } + "input": [{"role": "user", "content": "Calculate 123 times 456"}], + "reasoning": {"effort": "medium", "summary": "auto"}, + }, ) assert response.status_code == 200 @@ -620,10 +600,10 @@ class TestStewardIntegration: "input": [ {"role": "user", "content": "My favorite number is 42"}, {"role": "assistant", "content": "Noted, sir. 42 is an excellent choice."}, - {"role": "user", "content": "What was that number again?"} + {"role": "user", "content": "What was that number again?"}, ], - "reasoning": {"effort": "medium", "summary": "auto"} - } + "reasoning": {"effort": "medium", "summary": "auto"}, + }, ) assert response.status_code == 200 diff --git a/tests/e2e/test_orchestration_e2e.py b/tests/e2e/test_orchestration_e2e.py index 733c23b..50aab1c 100644 --- a/tests/e2e/test_orchestration_e2e.py +++ b/tests/e2e/test_orchestration_e2e.py @@ -18,13 +18,14 @@ Requirements: Note: LLM outputs are non-deterministic. Tests use flexible assertions that check for behavioral patterns rather than exact text matches. """ -import pytest -import httpx + import asyncio import re -from typing import AsyncGenerator +from collections.abc import AsyncGenerator from dataclasses import dataclass +import httpx +import pytest from src.core.config import PRODUCTION_TENANT, TEST_TENANT @@ -43,6 +44,7 @@ assert TEST_USER != PRODUCTION_TENANT @dataclass class LLMAssertionResult: """Result of an LLM output assertion check.""" + passed: bool evidence: str confidence: str # "high", "medium", "low" @@ -183,6 +185,7 @@ def assert_llm_behavior( # Fixtures # ============================================================================ + @pytest.fixture async def client() -> AsyncGenerator[httpx.AsyncClient, None]: """HTTP client for API requests.""" @@ -223,6 +226,7 @@ async def clean_test_memories(qdrant: QdrantVerifier): # Memory System Tests # ============================================================================ + @pytest.mark.e2e @pytest.mark.asyncio class TestMemoryStorage: @@ -245,9 +249,7 @@ class TestMemoryStorage: "/v1/responses", json={ "model": "Tatlock", - "input": [ - {"role": "user", "content": "Remember that my test color is purple"} - ], + "input": [{"role": "user", "content": "Remember that my test color is purple"}], }, ) @@ -281,7 +283,8 @@ class TestMemoryStorage: # Verify data in Qdrant - look for any color-related or test-related memory points = await qdrant.scroll_points(TEST_COLLECTION) relevant_memories = [ - p for p in points + p + for p in points if "color" in p.get("payload", {}).get("key", "").lower() or "purple" in str(p.get("payload", {}).get("value", "")).lower() or "test" in p.get("payload", {}).get("key", "").lower() @@ -341,9 +344,7 @@ class TestMemoryStorage: "/v1/responses", json={ "model": "Tatlock", - "input": [ - {"role": "user", "content": "What is my favorite food?"} - ], + "input": [{"role": "user", "content": "What is my favorite food?"}], }, ) @@ -431,6 +432,7 @@ class TestMemoryRecall: # Steward Delegation Tests # ============================================================================ + @pytest.mark.e2e @pytest.mark.asyncio class TestStewardDelegation: @@ -445,9 +447,7 @@ class TestStewardDelegation: "/v1/responses", json={ "model": "Tatlock", - "input": [ - {"role": "user", "content": "What do you know about me?"} - ], + "input": [{"role": "user", "content": "What do you know about me?"}], "reasoning": {"effort": "medium", "summary": "auto"}, }, ) @@ -479,9 +479,7 @@ class TestStewardDelegation: "/v1/responses", json={ "model": "Tatlock", - "input": [ - {"role": "user", "content": "Calculate 127 times 83"} - ], + "input": [{"role": "user", "content": "Calculate 127 times 83"}], "reasoning": {"effort": "medium", "summary": "auto"}, }, ) @@ -511,15 +509,16 @@ class TestStewardDelegation: # Remove commas for number comparison message_normalized = message_text.replace(",", "") - assert "10541" in message_normalized, ( - f"Expected calculation result 10541. Got: {message_text[:200]}" - ) + assert ( + "10541" in message_normalized + ), f"Expected calculation result 10541. Got: {message_text[:200]}" # ============================================================================ # Direct Delegation Tests # ============================================================================ + @pytest.mark.e2e @pytest.mark.asyncio class TestDirectDelegation: @@ -569,6 +568,7 @@ class TestDirectDelegation: # Data Verification Tests # ============================================================================ + @pytest.mark.e2e @pytest.mark.asyncio class TestDataVerification: @@ -610,8 +610,7 @@ class TestDataVerification: payload = point.get("payload", {}) for field in required_fields: assert field in payload, ( - f"Point {point['id']} missing required field '{field}'. " - f"Payload: {payload}" + f"Point {point['id']} missing required field '{field}'. " f"Payload: {payload}" ) async def test_qdrant_point_types_are_valid( @@ -643,6 +642,7 @@ class TestDataVerification: # Integration Health Tests # ============================================================================ + @pytest.mark.e2e @pytest.mark.asyncio class TestIntegrationHealth: @@ -678,6 +678,7 @@ class TestIntegrationHealth: # Orchestration Scenario Tests (from ORCHESTRATION_SCENARIOS.md) # ============================================================================ + @pytest.mark.e2e @pytest.mark.asyncio class TestScenario1WeatherWithMemory: @@ -697,9 +698,7 @@ class TestScenario1WeatherWithMemory: "/v1/responses", json={ "model": "Tatlock", - "input": [ - {"role": "user", "content": "What's the weather like?"} - ], + "input": [{"role": "user", "content": "What's the weather like?"}], "reasoning": {"effort": "medium", "summary": "auto"}, }, ) @@ -717,13 +716,17 @@ class TestScenario1WeatherWithMemory: print(f"Weather query - Steward: {reasoning_text[:200]}...") # Should mention location/memory and search capabilities - has_memory_mention = "biographer" in reasoning_text.lower() or "memory" in reasoning_text.lower() - has_search_mention = "tatlock_core" in reasoning_text.lower() or "search" in reasoning_text.lower() + has_memory_mention = ( + "biographer" in reasoning_text.lower() or "memory" in reasoning_text.lower() + ) + has_search_mention = ( + "tatlock_core" in reasoning_text.lower() or "search" in reasoning_text.lower() + ) # Weather query should trigger at least web search - assert has_search_mention, ( - f"Weather query should recommend search capability. Got: {reasoning_text[:200]}" - ) + assert ( + has_search_mention + ), f"Weather query should recommend search capability. Got: {reasoning_text[:200]}" @pytest.mark.e2e @@ -744,9 +747,7 @@ class TestScenario4SimpleExpertDelegation: "/v1/responses", json={ "model": "Tatlock", - "input": [ - {"role": "user", "content": "What is 847 times 293?"} - ], + "input": [{"role": "user", "content": "What is 847 times 293?"}], }, ) @@ -761,9 +762,9 @@ class TestScenario4SimpleExpertDelegation: # Expected result: 248171 (may be formatted as 248,171) message_normalized = message_text.replace(",", "") - assert "248171" in message_normalized, ( - f"Calculator should compute 847 * 293 = 248171. Got: {message_text[:200]}" - ) + assert ( + "248171" in message_normalized + ), f"Calculator should compute 847 * 293 = 248171. Got: {message_text[:200]}" async def test_datetime_query_uses_datetime_tool( self, @@ -774,9 +775,7 @@ class TestScenario4SimpleExpertDelegation: "/v1/responses", json={ "model": "Tatlock", - "input": [ - {"role": "user", "content": "What day of the week is it?"} - ], + "input": [{"role": "user", "content": "What day of the week is it?"}], }, ) @@ -791,8 +790,20 @@ class TestScenario4SimpleExpertDelegation: # Should mention a day of the week (full or abbreviated) days = [ - "monday", "tuesday", "wednesday", "thursday", "friday", "saturday", "sunday", - "mon", "tue", "wed", "thu", "fri", "sat", "sun" + "monday", + "tuesday", + "wednesday", + "thursday", + "friday", + "saturday", + "sunday", + "mon", + "tue", + "wed", + "thu", + "fri", + "sat", + "sun", ] has_day = any(day in message_text.lower() for day in days) @@ -838,9 +849,9 @@ class TestScenario6WikiCreation: print(f"Wiki search - Steward: {reasoning_text[:200]}...") # Should mention librarian - assert "librarian" in reasoning_text.lower(), ( - f"Wiki search should recommend librarian. Got: {reasoning_text[:200]}" - ) + assert ( + "librarian" in reasoning_text.lower() + ), f"Wiki search should recommend librarian. Got: {reasoning_text[:200]}" @pytest.mark.e2e @@ -862,7 +873,10 @@ class TestScenario8MultiExpertCoordination: json={ "model": "Tatlock", "input": [ - {"role": "user", "content": "What do you know about me? And also search for Python tutorials."} + { + "role": "user", + "content": "What do you know about me? And also search for Python tutorials.", + } ], "reasoning": {"effort": "medium", "summary": "auto"}, }, @@ -885,15 +899,16 @@ class TestScenario8MultiExpertCoordination: has_librarian = "librarian" in reasoning_text.lower() has_search = "tatlock_core" in reasoning_text.lower() or "search" in reasoning_text.lower() - assert has_biographer or has_librarian or has_search, ( - f"Complex query should identify multiple capabilities. Got: {reasoning_text[:200]}" - ) + assert ( + has_biographer or has_librarian or has_search + ), f"Complex query should identify multiple capabilities. Got: {reasoning_text[:200]}" # ============================================================================ # New Scenarios from Today's Session # ============================================================================ + @pytest.mark.e2e @pytest.mark.asyncio class TestDirectDelegationBypass: @@ -917,9 +932,7 @@ class TestDirectDelegationBypass: "/v1/responses", json={ "model": "Tatlock", - "input": [ - {"role": "user", "content": "Remember that my test value is alpha123"} - ], + "input": [{"role": "user", "content": "Remember that my test value is alpha123"}], }, ) @@ -979,7 +992,10 @@ class TestUserContextIsolation: json={ "model": "Tatlock", "input": [ - {"role": "user", "content": f"Remember that my isolation marker is {unique_value}"} + { + "role": "user", + "content": f"Remember that my isolation marker is {unique_value}", + } ], }, ) @@ -1000,13 +1016,15 @@ class TestUserContextIsolation: else: in_prod_collection = False - print(f"Isolation test - In test collection: {in_test_collection}, In prod: {in_prod_collection}") + print( + f"Isolation test - In test collection: {in_test_collection}, In prod: {in_prod_collection}" + ) # Should be in test collection OR response acknowledged # Should NOT be in production collection - assert not in_prod_collection, ( - f"Test data leaked to production collection! Value: {unique_value}" - ) + assert ( + not in_prod_collection + ), f"Test data leaked to production collection! Value: {unique_value}" @pytest.mark.e2e @@ -1025,16 +1043,16 @@ class TestErrorHandling: "/v1/responses", json={ "model": "nonexistent-model-xyz", - "input": [ - {"role": "user", "content": "Hello"} - ], + "input": [{"role": "user", "content": "Hello"}], }, ) # Should return error status (404 or 400) - assert response.status_code in [400, 404, 422], ( - f"Expected error status for invalid model. Got: {response.status_code}" - ) + assert response.status_code in [ + 400, + 404, + 422, + ], f"Expected error status for invalid model. Got: {response.status_code}" data = response.json() # Error could be in "error" or "detail" key assert "error" in data or "detail" in data @@ -1084,7 +1102,10 @@ class TestEvaluationReport: json={ "model": "Tatlock", "input": [ - {"role": "user", "content": "Remember that my test pet is a hamster named Fluffy"} + { + "role": "user", + "content": "Remember that my test pet is a hamster named Fluffy", + } ], }, ) @@ -1096,7 +1117,7 @@ class TestEvaluationReport: break report_lines.append("\n[STORE TEST]") - report_lines.append(f"Input: 'Remember that my test pet is a hamster named Fluffy'") + report_lines.append("Input: 'Remember that my test pet is a hamster named Fluffy'") report_lines.append(f"Response: {store_text[:200]}...") await asyncio.sleep(1.0) @@ -1106,9 +1127,7 @@ class TestEvaluationReport: "/v1/responses", json={ "model": "Tatlock", - "input": [ - {"role": "user", "content": "What pet do I have?"} - ], + "input": [{"role": "user", "content": "What pet do I have?"}], }, ) @@ -1119,7 +1138,7 @@ class TestEvaluationReport: break report_lines.append("\n[RECALL TEST]") - report_lines.append(f"Input: 'What pet do I have?'") + report_lines.append("Input: 'What pet do I have?'") report_lines.append(f"Response: {recall_text[:200]}...") # Check Qdrant state @@ -1128,7 +1147,8 @@ class TestEvaluationReport: report_lines.append(f"Total points in {TEST_COLLECTION}: {len(points)}") pet_memories = [ - p for p in points + p + for p in points if "pet" in str(p.get("payload", {})).lower() or "fluffy" in str(p.get("payload", {})).lower() or "hamster" in str(p.get("payload", {})).lower() diff --git a/tests/integration/test_steward_streaming.py b/tests/integration/test_steward_streaming.py index aff5050..57bdc91 100644 --- a/tests/integration/test_steward_streaming.py +++ b/tests/integration/test_steward_streaming.py @@ -3,12 +3,14 @@ Integration tests for Steward + Tatlock streaming. Tests the complete streaming flow with Steward preprocessing. """ -import pytest -from unittest.mock import AsyncMock, MagicMock, patch -from src.responses.schemas import ResponseRequest -from src.responses.streaming import StreamingCoordinator, StreamEventType +from unittest.mock import patch + +import pytest + from src.core.startup import initialize_application +from src.responses.schemas import ResponseRequest +from src.responses.streaming import StreamEventType, StreamingCoordinator @pytest.fixture(scope="module", autouse=True) @@ -32,7 +34,9 @@ class TestStewardStreaming: # Mock the Steward analysis with patch("src.core.preprocessing.analyze_request") as mock_steward: # Mock the streaming method (async generator) - with patch("src.agents.tatlock.TatlockAgent.run_with_scoped_tools_stream") as mock_tatlock_stream: + with patch( + "src.agents.tatlock.TatlockAgent.run_with_scoped_tools_stream" + ) as mock_tatlock_stream: from src.agents.steward.schemas import ConversationContext, StewardRecommendation # Mock Steward recommendation @@ -89,7 +93,9 @@ class TestStewardStreaming: ) with patch("src.core.preprocessing.analyze_request") as mock_steward: - with patch("src.agents.tatlock.TatlockAgent.run_with_scoped_tools_stream") as mock_tatlock_stream: + with patch( + "src.agents.tatlock.TatlockAgent.run_with_scoped_tools_stream" + ) as mock_tatlock_stream: from src.agents.steward.schemas import ConversationContext, StewardRecommendation mock_steward.return_value = StewardRecommendation( @@ -99,7 +105,7 @@ class TestStewardStreaming: conversation_context=ConversationContext( has_previous_context=True, relevant_turns=[0], - context_summary="Previous calculation in turn 0" + context_summary="Previous calculation in turn 0", ), ) @@ -134,7 +140,9 @@ class TestStewardStreaming: ) with patch("src.core.preprocessing.analyze_request") as mock_steward: - with patch("src.agents.tatlock.TatlockAgent.run_with_scoped_tools_stream") as mock_tatlock_stream: + with patch( + "src.agents.tatlock.TatlockAgent.run_with_scoped_tools_stream" + ) as mock_tatlock_stream: from src.agents.steward.schemas import ConversationContext, StewardRecommendation mock_steward.return_value = StewardRecommendation( @@ -173,7 +181,9 @@ class TestStewardStreaming: ) with patch("src.core.preprocessing.analyze_request") as mock_steward: - with patch("src.agents.tatlock.TatlockAgent.run_with_scoped_tools_stream") as mock_tatlock_stream: + with patch( + "src.agents.tatlock.TatlockAgent.run_with_scoped_tools_stream" + ) as mock_tatlock_stream: from src.agents.steward.schemas import ConversationContext, StewardRecommendation mock_steward.return_value = StewardRecommendation( diff --git a/tests/integration/test_steward_tatlock_integration.py b/tests/integration/test_steward_tatlock_integration.py index 903516f..bd45e3a 100644 --- a/tests/integration/test_steward_tatlock_integration.py +++ b/tests/integration/test_steward_tatlock_integration.py @@ -7,12 +7,14 @@ Tests the complete Phase 2 request pipeline: 3. Tatlock runs with scoped tools 4. Response includes both Steward reasoning and Tatlock output """ -import pytest + from unittest.mock import AsyncMock, MagicMock, patch +import pytest + +from src.core.startup import initialize_application from src.responses.schemas import ResponseRequest from src.responses.service import create_response_with_steward -from src.core.startup import initialize_application @pytest.fixture(scope="module", autouse=True) @@ -56,8 +58,9 @@ class TestStewardTatlockIntegration: assert mock_steward.called # Note: preprocess_request injects temporal context steward_call_arg = mock_steward.call_args[0][0] - assert steward_call_arg.startswith("What's 2 + 2?"), \ - f"Expected request to start with original message, got: {steward_call_arg}" + assert steward_call_arg.startswith( + "What's 2 + 2?" + ), f"Expected request to start with original message, got: {steward_call_arg}" # Verify Tatlock was called with scoped tools assert mock_tatlock.called @@ -100,7 +103,7 @@ class TestStewardTatlockIntegration: conversation_context=ConversationContext( has_previous_context=True, relevant_turns=[0], - context_summary="Previous calculation in turn 0" + context_summary="Previous calculation in turn 0", ), ) @@ -161,7 +164,10 @@ class TestStewardTatlockIntegration: with patch("src.core.preprocessing.analyze_request") as mock_steward: with patch("src.agents.tatlock.TatlockAgent.run_with_scoped_tools") as mock_tatlock: with patch("src.core.tool_tracking.ToolCallTracker.finalize") as mock_finalize: - from src.agents.steward.schemas import ConversationContext, StewardRecommendation + from src.agents.steward.schemas import ( + ConversationContext, + StewardRecommendation, + ) mock_steward.return_value = StewardRecommendation( recommended_capabilities=["tatlock_core"], @@ -198,7 +204,9 @@ class TestStewardTatlockIntegration: missing_capabilities="Image generation capability would be needed", ) - mock_tatlock.return_value = "I'm afraid I don't have image generation capabilities, sir." + mock_tatlock.return_value = ( + "I'm afraid I don't have image generation capabilities, sir." + ) response = await create_response_with_steward(request) @@ -220,7 +228,10 @@ class TestStewardTatlockIntegration: with patch("src.core.preprocessing.analyze_request") as mock_steward: with patch("src.agents.tatlock.TatlockAgent.run_with_scoped_tools") as mock_tatlock: with patch("src.responses.service.ToolCallTracker") as mock_tracker_class: - from src.agents.steward.schemas import ConversationContext, StewardRecommendation + from src.agents.steward.schemas import ( + ConversationContext, + StewardRecommendation, + ) mock_steward.return_value = StewardRecommendation( recommended_capabilities=["tatlock_core"], diff --git a/tests/integration/test_tatlock_streaming.py b/tests/integration/test_tatlock_streaming.py index 0d80eb5..52eba10 100644 --- a/tests/integration/test_tatlock_streaming.py +++ b/tests/integration/test_tatlock_streaming.py @@ -5,10 +5,12 @@ These tests verify the complete streaming flow from API endpoint through StreamingCoordinator to TatlockAgent, ensuring no text duplication and proper delta calculation. """ + import json + import pytest -from httpx import AsyncClient from fastapi.testclient import TestClient +from httpx import AsyncClient @pytest.mark.integration @@ -24,7 +26,7 @@ async def test_tatlock_streaming_no_duplication(async_client: AsyncClient): request_data = { "model": "Tatlock", "input": [{"role": "user", "content": "Say hello"}], - "stream": True + "stream": True, } collected_deltas = [] @@ -74,12 +76,12 @@ async def test_tatlock_streaming_no_duplication(async_client: AsyncClient): if len(words) > 0: # Check for consecutive duplicate words (sign of duplication bug) consecutive_dupes = sum( - 1 for i in range(len(words) - 1) - if words[i] == words[i + 1] and len(words[i]) > 3 + 1 for i in range(len(words) - 1) if words[i] == words[i + 1] and len(words[i]) > 3 ) # Allow a few duplicates (natural language), but not excessive - assert consecutive_dupes < len(words) * 0.1, \ - f"Too many consecutive duplicate words: {consecutive_dupes}/{len(words)}" + assert ( + consecutive_dupes < len(words) * 0.1 + ), f"Too many consecutive duplicate words: {consecutive_dupes}/{len(words)}" @pytest.mark.integration @@ -94,7 +96,7 @@ async def test_tatlock_chat_streaming_no_duplication(async_client: AsyncClient): request_data = { "model": "Tatlock", "messages": [{"role": "user", "content": "Hello"}], - "stream": True + "stream": True, } collected_content = [] @@ -138,11 +140,11 @@ async def test_tatlock_chat_streaming_no_duplication(async_client: AsyncClient): words = full_response.lower().split() if len(words) > 0: consecutive_dupes = sum( - 1 for i in range(len(words) - 1) - if words[i] == words[i + 1] and len(words[i]) > 3 + 1 for i in range(len(words) - 1) if words[i] == words[i + 1] and len(words[i]) > 3 ) - assert consecutive_dupes < len(words) * 0.1, \ - f"Too many consecutive duplicate words in chat response: {consecutive_dupes}/{len(words)}" + assert ( + consecutive_dupes < len(words) * 0.1 + ), f"Too many consecutive duplicate words in chat response: {consecutive_dupes}/{len(words)}" @pytest.mark.integration @@ -153,7 +155,7 @@ def test_tatlock_non_streaming_responses_api(client: TestClient): request_data = { "model": "Tatlock", "input": [{"role": "user", "content": "Say hello"}], - "stream": False + "stream": False, } response = client.post("/v1/responses", json=request_data, timeout=30.0) @@ -183,7 +185,7 @@ def test_tatlock_non_streaming_chat_api(client: TestClient): request_data = { "model": "Tatlock", "messages": [{"role": "user", "content": "Hello"}], - "stream": False + "stream": False, } response = client.post("/v1/chat/completions", json=request_data, timeout=30.0) @@ -216,7 +218,7 @@ async def test_tatlock_streaming_delta_accumulation(async_client: AsyncClient): request_data = { "model": "Tatlock", "input": [{"role": "user", "content": "Count to three"}], - "stream": True + "stream": True, } collected_deltas = [] @@ -248,8 +250,9 @@ async def test_tatlock_streaming_delta_accumulation(async_client: AsyncClient): # Verify each delta is new content current_full = "".join(collected_deltas) - assert current_full.startswith(previous_full_text), \ - "Deltas should accumulate progressively" + assert current_full.startswith( + previous_full_text + ), "Deltas should accumulate progressively" previous_full_text = current_full except json.JSONDecodeError: @@ -273,7 +276,7 @@ async def test_tatlock_with_reasoning(async_client: AsyncClient): "model": "Tatlock", "input": [{"role": "user", "content": "Hello"}], "reasoning": {"effort": "medium", "summary": "auto"}, - "stream": True + "stream": True, } has_reasoning = False @@ -328,7 +331,7 @@ async def test_tatlock_markdown_formatting_preserved(async_client: AsyncClient): request_data = { "model": "Tatlock", "input": [{"role": "user", "content": "Can you give me an HTML5 boilerplate template?"}], - "stream": True + "stream": True, } collected_deltas = [] @@ -365,15 +368,15 @@ async def test_tatlock_markdown_formatting_preserved(async_client: AsyncClient): full_response = "".join(collected_deltas) # Always print the response for debugging - print("\n" + "="*80) + print("\n" + "=" * 80) print("FULL RESPONSE (repr):") - print("="*80) + print("=" * 80) print(repr(full_response)) - print("\n" + "="*80) + print("\n" + "=" * 80) print("FULL RESPONSE (formatted):") - print("="*80) + print("=" * 80) print(full_response) - print("="*80 + "\n") + print("=" * 80 + "\n") # Verify we got a response (xfail if LLM didn't produce output) if len(full_response) < 100: @@ -384,7 +387,7 @@ async def test_tatlock_markdown_formatting_preserved(async_client: AsyncClient): pytest.xfail("No markdown code blocks in response (LLM response varied)") # Verify newlines are preserved (not all collapsed to spaces) - newline_count = full_response.count('\n') + newline_count = full_response.count("\n") if newline_count < 5: pytest.xfail(f"Only {newline_count} newlines, formatting may have been lost") @@ -412,7 +415,7 @@ def test_tatlock_markdown_non_streaming(client: TestClient): request_data = { "model": "Tatlock", "input": [{"role": "user", "content": "Give me a simple Python hello world code"}], - "stream": False + "stream": False, } response = client.post("/v1/responses", json=request_data, timeout=30.0) diff --git a/tests/models/test_router.py b/tests/models/test_router.py index bfe5182..1614e30 100644 --- a/tests/models/test_router.py +++ b/tests/models/test_router.py @@ -1,6 +1,7 @@ """ Tests for models listing router. """ + import pytest from fastapi.testclient import TestClient @@ -39,10 +40,10 @@ def test_list_models(client: TestClient) -> None: def test_models_endpoint_returns_json(client: TestClient) -> None: """Test that models endpoint returns valid JSON.""" response = client.get("/v1/models") - + assert response.status_code == 200 assert response.headers["content-type"] == "application/json" - + # Should be able to parse as JSON data = response.json() assert isinstance(data, dict) diff --git a/tests/responses/test_advanced_features.py b/tests/responses/test_advanced_features.py index 6c77e2d..5de876d 100644 --- a/tests/responses/test_advanced_features.py +++ b/tests/responses/test_advanced_features.py @@ -6,7 +6,9 @@ Tests: - Stop sequence detection and enforcement - Max tokens enforcement """ + import json + import pytest from fastapi.testclient import TestClient from httpx import AsyncClient @@ -15,11 +17,11 @@ from pydantic import ValidationError from src.responses.schemas import ResponseRequest from src.responses.streaming import StreamingCoordinator - # ============================================================================ # Parameter Validation Tests # ============================================================================ + @pytest.mark.unit def test_temperature_validation(): """Test temperature parameter validation.""" @@ -27,9 +29,7 @@ def test_temperature_validation(): valid_temps = [0.0, 0.5, 1.0, 1.5, 2.0] for temp in valid_temps: request = ResponseRequest( - model="lorem-tester", - input=[{"role": "user", "content": "Hello"}], - temperature=temp + model="lorem-tester", input=[{"role": "user", "content": "Hello"}], temperature=temp ) assert request.temperature == temp @@ -38,7 +38,7 @@ def test_temperature_validation(): ResponseRequest( model="lorem-tester", input=[{"role": "user", "content": "Hello"}], - temperature=-0.1 # Too low + temperature=-0.1, # Too low ) assert "temperature" in str(exc_info.value).lower() @@ -46,7 +46,7 @@ def test_temperature_validation(): ResponseRequest( model="lorem-tester", input=[{"role": "user", "content": "Hello"}], - temperature=2.1 # Too high + temperature=2.1, # Too high ) assert "temperature" in str(exc_info.value).lower() @@ -55,12 +55,12 @@ def test_temperature_validation(): def test_reasoning_effort_validation(): """Test reasoning.effort parameter validation.""" # Valid effort levels - valid_efforts = ['none', 'minimal', 'low', 'medium', 'high', 'xhigh'] + valid_efforts = ["none", "minimal", "low", "medium", "high", "xhigh"] for effort in valid_efforts: request = ResponseRequest( model="lorem-tester", input=[{"role": "user", "content": "Hello"}], - reasoning={"effort": effort, "summary": "auto"} + reasoning={"effort": effort, "summary": "auto"}, ) assert request.reasoning["effort"] == effort @@ -69,7 +69,7 @@ def test_reasoning_effort_validation(): ResponseRequest( model="lorem-tester", input=[{"role": "user", "content": "Hello"}], - reasoning={"effort": "invalid", "summary": "auto"} + reasoning={"effort": "invalid", "summary": "auto"}, ) assert "reasoning.effort" in str(exc_info.value) @@ -78,12 +78,12 @@ def test_reasoning_effort_validation(): def test_reasoning_summary_validation(): """Test reasoning.summary parameter validation.""" # Valid summary values - valid_summaries = ['auto', 'off'] + valid_summaries = ["auto", "off"] for summary in valid_summaries: request = ResponseRequest( model="lorem-tester", input=[{"role": "user", "content": "Hello"}], - reasoning={"effort": "medium", "summary": summary} + reasoning={"effort": "medium", "summary": summary}, ) assert request.reasoning["summary"] == summary @@ -92,7 +92,7 @@ def test_reasoning_summary_validation(): ResponseRequest( model="lorem-tester", input=[{"role": "user", "content": "Hello"}], - reasoning={"effort": "medium", "summary": "invalid"} + reasoning={"effort": "medium", "summary": "invalid"}, ) assert "reasoning.summary" in str(exc_info.value) @@ -102,26 +102,20 @@ def test_max_output_tokens_validation(): """Test max_output_tokens parameter validation.""" # Valid values request = ResponseRequest( - model="lorem-tester", - input=[{"role": "user", "content": "Hello"}], - max_output_tokens=100 + model="lorem-tester", input=[{"role": "user", "content": "Hello"}], max_output_tokens=100 ) assert request.max_output_tokens == 100 # None is valid (unlimited) request = ResponseRequest( - model="lorem-tester", - input=[{"role": "user", "content": "Hello"}], - max_output_tokens=None + model="lorem-tester", input=[{"role": "user", "content": "Hello"}], max_output_tokens=None ) assert request.max_output_tokens is None # Invalid: zero or negative with pytest.raises(ValidationError) as exc_info: ResponseRequest( - model="lorem-tester", - input=[{"role": "user", "content": "Hello"}], - max_output_tokens=0 + model="lorem-tester", input=[{"role": "user", "content": "Hello"}], max_output_tokens=0 ) assert "max_output_tokens" in str(exc_info.value) @@ -129,7 +123,7 @@ def test_max_output_tokens_validation(): ResponseRequest( model="lorem-tester", input=[{"role": "user", "content": "Hello"}], - max_output_tokens=-10 + max_output_tokens=-10, ) assert "max_output_tokens" in str(exc_info.value) @@ -141,9 +135,7 @@ def test_stop_sequences_validation(): for num_seqs in range(1, 5): stop_seqs = [f"stop{i}" for i in range(num_seqs)] request = ResponseRequest( - model="lorem-tester", - input=[{"role": "user", "content": "Hello"}], - stop=stop_seqs + model="lorem-tester", input=[{"role": "user", "content": "Hello"}], stop=stop_seqs ) assert request.stop == stop_seqs @@ -152,7 +144,7 @@ def test_stop_sequences_validation(): ResponseRequest( model="lorem-tester", input=[{"role": "user", "content": "Hello"}], - stop=["stop1", "stop2", "stop3", "stop4", "stop5"] # 5 sequences + stop=["stop1", "stop2", "stop3", "stop4", "stop5"], # 5 sequences ) assert "4 stop sequences" in str(exc_info.value) @@ -161,7 +153,7 @@ def test_stop_sequences_validation(): ResponseRequest( model="lorem-tester", input=[{"role": "user", "content": "Hello"}], - stop=["stop1", ""] # Empty string + stop=["stop1", ""], # Empty string ) assert "non-empty" in str(exc_info.value).lower() @@ -170,6 +162,7 @@ def test_stop_sequences_validation(): # Stop Sequence Enforcement Tests # ============================================================================ + @pytest.mark.unit def test_stop_sequence_detection_helper(): """Test stop sequence detection helper method.""" @@ -181,26 +174,17 @@ def test_stop_sequence_detection_helper(): assert text == "Hello world" # Stop sequence not present - found, text = coordinator._check_stop_sequence( - "Hello world", - ["STOP", "END"] - ) + found, text = coordinator._check_stop_sequence("Hello world", ["STOP", "END"]) assert found is False assert text == "Hello world" # Stop sequence found - found, text = coordinator._check_stop_sequence( - "Hello STOP this should not appear", - ["STOP"] - ) + found, text = coordinator._check_stop_sequence("Hello STOP this should not appear", ["STOP"]) assert found is True assert text == "Hello " # Multiple stop sequences, first one wins - found, text = coordinator._check_stop_sequence( - "Hello STOP this END that", - ["STOP", "END"] - ) + found, text = coordinator._check_stop_sequence("Hello STOP this END that", ["STOP", "END"]) assert found is True assert text == "Hello " @@ -216,7 +200,7 @@ async def test_stop_sequence_in_streaming(async_client: AsyncClient): "model": "lorem-tester", "input": [{"role": "user", "content": "Generate long text"}], "stop": ["dolor"], # Common word in lorem ipsum - "stream": True + "stream": True, } chunks_received = [] @@ -251,6 +235,7 @@ async def test_stop_sequence_in_streaming(async_client: AsyncClient): # Max Tokens Enforcement Tests # ============================================================================ + @pytest.mark.unit def test_max_tokens_check_helper(): """Test max tokens check helper method.""" @@ -296,7 +281,7 @@ async def test_max_tokens_in_streaming(async_client: AsyncClient): "model": "lorem-tester", "input": [{"role": "user", "content": "Generate text"}], "max_output_tokens": 5, # Very low limit - "stream": True + "stream": True, } chunks_received = [] @@ -340,6 +325,7 @@ async def test_max_tokens_in_streaming(async_client: AsyncClient): # Combined Features Test # ============================================================================ + @pytest.mark.unit def test_combined_validation(client: TestClient): """Test combined parameter validation in actual request.""" @@ -351,7 +337,7 @@ def test_combined_validation(client: TestClient): "max_output_tokens": 100, "stop": ["STOP", "END"], "reasoning": {"effort": "high", "summary": "auto"}, - "stream": False + "stream": False, } response = client.post("/v1/responses", json=request_data) @@ -368,7 +354,7 @@ def test_invalid_combined_parameters(client: TestClient): "model": "lorem-tester", "input": [{"role": "user", "content": "Hello"}], "temperature": 3.0, # Too high - "stream": False + "stream": False, } response = client.post("/v1/responses", json=request_data) @@ -382,6 +368,7 @@ def test_invalid_combined_parameters(client: TestClient): # Streaming Delta Calculation Tests (No Duplication) # ============================================================================ + @pytest.mark.unit @pytest.mark.asyncio async def test_streaming_delta_calculation_no_duplication(): @@ -392,8 +379,10 @@ async def test_streaming_delta_calculation_no_duplication(): This test prevents the duplication bug where the same text was streamed multiple times because we weren't computing deltas correctly. """ + from collections.abc import AsyncGenerator + from typing import Any + from src.agents.base import AgentInterface, OutputItem - from typing import AsyncGenerator, Any # Create a mock agent that simulates PydanticAI's behavior # (yielding accumulated text, not deltas) @@ -406,7 +395,7 @@ async def test_streaming_delta_calculation_no_duplication(): temperature: float = 1.0, max_tokens: int | None = None, stop: list[str] | None = None, - **kwargs: Any + **kwargs: Any, ) -> AsyncGenerator[OutputItem, None]: """ Simulate PydanticAI streaming behavior: @@ -430,12 +419,8 @@ async def test_streaming_delta_calculation_no_duplication(): type="message", id=msg_id, role="assistant", - content=[{ - "type": "output_text", - "text": text, - "annotations": [] - }], - status="in_progress" + content=[{"type": "output_text", "text": text, "annotations": []}], + status="in_progress", ) # Final message @@ -443,12 +428,10 @@ async def test_streaming_delta_calculation_no_duplication(): type="message", id=msg_id, role="assistant", - content=[{ - "type": "output_text", - "text": "Hello world how are you", - "annotations": [] - }], - status="completed" + content=[ + {"type": "output_text", "text": "Hello world how are you", "annotations": []} + ], + status="completed", ) async def supports_tools(self) -> bool: @@ -462,7 +445,9 @@ async def test_streaming_delta_calculation_no_duplication(): # Register the mock agent import time + from src.agents.registry import ModelRegistry + ModelRegistry.MODELS["mock-streaming"] = { "agent_class": MockStreamingAgent, "description": "Mock streaming agent for testing", @@ -473,9 +458,7 @@ async def test_streaming_delta_calculation_no_duplication(): try: # Create a test request request = ResponseRequest( - model="mock-streaming", - input=[{"role": "user", "content": "Test"}], - stream=True + model="mock-streaming", input=[{"role": "user", "content": "Test"}], stream=True ) # Stream the response @@ -512,8 +495,10 @@ async def test_streaming_with_multiple_message_items(): Test that coordinator handles multiple message OutputItems correctly, only streaming the delta between each one. """ + from collections.abc import AsyncGenerator + from typing import Any + from src.agents.base import AgentInterface, OutputItem - from typing import AsyncGenerator, Any class MockMultiMessageAgent(AgentInterface): async def generate_response( @@ -524,7 +509,7 @@ async def test_streaming_with_multiple_message_items(): temperature: float = 1.0, max_tokens: int | None = None, stop: list[str] | None = None, - **kwargs: Any + **kwargs: Any, ) -> AsyncGenerator[OutputItem, None]: """Yield multiple in_progress messages with accumulated text.""" # First chunk @@ -533,7 +518,7 @@ async def test_streaming_with_multiple_message_items(): id="msg_1", role="assistant", content=[{"type": "output_text", "text": "The answer is", "annotations": []}], - status="in_progress" + status="in_progress", ) # Second chunk (more text accumulated) @@ -542,7 +527,7 @@ async def test_streaming_with_multiple_message_items(): id="msg_1", role="assistant", content=[{"type": "output_text", "text": "The answer is 42", "annotations": []}], - status="in_progress" + status="in_progress", ) # Final chunk @@ -551,7 +536,7 @@ async def test_streaming_with_multiple_message_items(): id="msg_1", role="assistant", content=[{"type": "output_text", "text": "The answer is 42", "annotations": []}], - status="completed" + status="completed", ) async def supports_tools(self) -> bool: @@ -565,7 +550,9 @@ async def test_streaming_with_multiple_message_items(): # Register mock agent import time + from src.agents.registry import ModelRegistry + ModelRegistry.MODELS["mock-multi"] = { "agent_class": MockMultiMessageAgent, "description": "Mock multi-message agent for testing", @@ -577,7 +564,7 @@ async def test_streaming_with_multiple_message_items(): request = ResponseRequest( model="mock-multi", input=[{"role": "user", "content": "What is the answer?"}], - stream=True + stream=True, ) coordinator = StreamingCoordinator() diff --git a/tests/responses/test_error_handling.py b/tests/responses/test_error_handling.py index 20515ff..5c29616 100644 --- a/tests/responses/test_error_handling.py +++ b/tests/responses/test_error_handling.py @@ -1,7 +1,9 @@ """ Tests for error handling in Responses API. """ + import json + import pytest from fastapi.testclient import TestClient from httpx import AsyncClient @@ -12,10 +14,8 @@ def test_model_not_found_error(client: TestClient) -> None: """Test response when model doesn't exist.""" request_data = { "model": "nonexistent-model-12345", - "input": [ - {"role": "user", "content": "Hello"} - ], - "stream": False + "input": [{"role": "user", "content": "Hello"}], + "stream": False, } response = client.post("/v1/responses", json=request_data) @@ -29,11 +29,7 @@ def test_model_not_found_error(client: TestClient) -> None: @pytest.mark.unit def test_validation_error_missing_model(client: TestClient) -> None: """Test validation error when model field is missing.""" - request_data = { - "input": [ - {"role": "user", "content": "Hello"} - ] - } + request_data = {"input": [{"role": "user", "content": "Hello"}]} response = client.post("/v1/responses", json=request_data) @@ -45,10 +41,8 @@ def test_validation_error_invalid_temperature(client: TestClient) -> None: """Test validation error for out-of-range temperature.""" request_data = { "model": "lorem-tester", - "input": [ - {"role": "user", "content": "Hello"} - ], - "temperature": 3.0 # Max is 2.0 + "input": [{"role": "user", "content": "Hello"}], + "temperature": 3.0, # Max is 2.0 } response = client.post("/v1/responses", json=request_data) @@ -61,10 +55,8 @@ def test_rate_limit_error(client: TestClient) -> None: """Test rate limit error trigger.""" request_data = { "model": "lorem-tester", - "input": [ - {"role": "user", "content": "trigger_rate_limit"} - ], - "stream": False + "input": [{"role": "user", "content": "trigger_rate_limit"}], + "stream": False, } response = client.post("/v1/responses", json=request_data) @@ -78,10 +70,8 @@ def test_context_overflow_error(client: TestClient) -> None: """Test context length overflow error trigger.""" request_data = { "model": "lorem-tester", - "input": [ - {"role": "user", "content": "trigger_context_overflow"} - ], - "stream": False + "input": [{"role": "user", "content": "trigger_context_overflow"}], + "stream": False, } response = client.post("/v1/responses", json=request_data) @@ -96,10 +86,8 @@ async def test_streaming_model_not_found(async_client: AsyncClient) -> None: """Test streaming response with nonexistent model.""" request_data = { "model": "nonexistent-streaming-model", - "input": [ - {"role": "user", "content": "Hello"} - ], - "stream": True + "input": [{"role": "user", "content": "Hello"}], + "stream": True, } async with async_client.stream( @@ -133,10 +121,8 @@ async def test_streaming_rate_limit_error(async_client: AsyncClient) -> None: """Test streaming with rate limit error.""" request_data = { "model": "lorem-tester", - "input": [ - {"role": "user", "content": "trigger_rate_limit"} - ], - "stream": True + "input": [{"role": "user", "content": "trigger_rate_limit"}], + "stream": True, } async with async_client.stream( @@ -177,10 +163,8 @@ async def test_streaming_context_overflow_error(async_client: AsyncClient) -> No """Test streaming with context overflow error.""" request_data = { "model": "lorem-tester", - "input": [ - {"role": "user", "content": "trigger_context_overflow"} - ], - "stream": True + "input": [{"role": "user", "content": "trigger_context_overflow"}], + "stream": True, } async with async_client.stream( @@ -215,17 +199,9 @@ def test_function_call_output_item(client: TestClient) -> None: """Test response with function call items.""" request_data = { "model": "lorem-tester", - "input": [ - {"role": "user", "content": "Use the search tool"} - ], - "tools": [ - { - "name": "search", - "description": "Search for information", - "parameters": {} - } - ], - "stream": False + "input": [{"role": "user", "content": "Use the search tool"}], + "tools": [{"name": "search", "description": "Search for information", "parameters": {}}], + "stream": False, } response = client.post("/v1/responses", json=request_data) @@ -248,10 +224,8 @@ def test_pipeline_prefix_stripping(client: TestClient) -> None: """Test that pipeline prefixes are stripped from model names.""" request_data = { "model": "some_pipeline.lorem-tester", - "input": [ - {"role": "user", "content": "Hello"} - ], - "stream": False + "input": [{"role": "user", "content": "Hello"}], + "stream": False, } response = client.post("/v1/responses", json=request_data) @@ -267,10 +241,8 @@ def test_multiple_pipeline_prefixes(client: TestClient) -> None: """Test multiple dots in model name (only first is prefix).""" request_data = { "model": "pipeline.sub.lorem-tester", - "input": [ - {"role": "user", "content": "Hello"} - ], - "stream": False + "input": [{"role": "user", "content": "Hello"}], + "stream": False, } response = client.post("/v1/responses", json=request_data) diff --git a/tests/responses/test_history.py b/tests/responses/test_history.py index 50413b8..3a616a0 100644 --- a/tests/responses/test_history.py +++ b/tests/responses/test_history.py @@ -1,13 +1,14 @@ """ Tests for conversation history management. """ + import pytest from fastapi.testclient import TestClient -from src.responses.history import ConversationHistory -from src.responses.context import ContextWindow -from src.responses.schemas import ResponseRequest from src.responses import service +from src.responses.context import ContextWindow +from src.responses.history import ConversationHistory +from src.responses.schemas import ResponseRequest @pytest.mark.unit @@ -19,7 +20,7 @@ async def test_conversation_id_from_metadata(): request = ResponseRequest( model="lorem-tester", input=[{"role": "user", "content": "Hello"}], - metadata={"conversation_id": "conv_123"} + metadata={"conversation_id": "conv_123"}, ) conv_id = await history.get_conversation_id(request) @@ -34,7 +35,7 @@ async def test_conversation_id_generation(): request = ResponseRequest( model="lorem-tester", - input=[{"role": "user", "content": "Hello"}] + input=[{"role": "user", "content": "Hello"}], # No metadata provided ) @@ -43,10 +44,7 @@ async def test_conversation_id_generation(): assert len(conv_id) == 16 # 16 character hex # Same first message should generate same ID - request2 = ResponseRequest( - model="lorem-tester", - input=[{"role": "user", "content": "Hello"}] - ) + request2 = ResponseRequest(model="lorem-tester", input=[{"role": "user", "content": "Hello"}]) conv_id2 = await history.get_conversation_id(request2) assert conv_id == conv_id2 @@ -59,7 +57,7 @@ def test_conversation_history_tracking(client: TestClient): "model": "lorem-tester", "input": [{"role": "user", "content": "Hello"}], "metadata": {"conversation_id": "test_conv_001"}, - "stream": False + "stream": False, } response = client.post("/v1/responses", json=request_data) @@ -71,10 +69,10 @@ def test_conversation_history_tracking(client: TestClient): "input": [ {"role": "user", "content": "Hello"}, {"role": "assistant", "content": "Hi there!"}, - {"role": "user", "content": "How are you?"} + {"role": "user", "content": "How are you?"}, ], "metadata": {"conversation_id": "test_conv_001"}, - "stream": False + "stream": False, } response2 = client.post("/v1/responses", json=request_data2) @@ -94,7 +92,7 @@ async def test_conversation_history_retrieval(): request = ResponseRequest( model="lorem-tester", input=[{"role": "user", "content": "Test"}], - metadata={"conversation_id": "test_retrieve"} + metadata={"conversation_id": "test_retrieve"}, ) conv_id = await history.get_conversation_id(request) @@ -104,24 +102,15 @@ async def test_conversation_history_retrieval(): assert len(retrieved) == 0 # Add mock response - from src.responses.schemas import Response, ResponseUsage, MessageOutputItem, OutputTextContent + from src.responses.schemas import MessageOutputItem, OutputTextContent, Response, ResponseUsage + mock_response = Response( id="resp_123", created_at=1234567890, model="lorem-tester", status="completed", - output=[ - MessageOutputItem( - id="msg_1", - content=[OutputTextContent(text="Test response")] - ) - ], - usage=ResponseUsage( - input_tokens=10, - output_tokens=5, - reasoning_tokens=0, - total_tokens=15 - ) + output=[MessageOutputItem(id="msg_1", content=[OutputTextContent(text="Test response")])], + usage=ResponseUsage(input_tokens=10, output_tokens=5, reasoning_tokens=0, total_tokens=15), ) await history.add_response(conv_id, mock_response) @@ -138,7 +127,7 @@ async def test_conversation_history_trimming(): """Test that history is trimmed to max_turns.""" history = ConversationHistory(max_turns=3) - from src.responses.schemas import Response, ResponseUsage, MessageOutputItem, OutputTextContent + from src.responses.schemas import MessageOutputItem, OutputTextContent, Response, ResponseUsage conv_id = "test_trim" @@ -150,17 +139,11 @@ async def test_conversation_history_trimming(): model="lorem-tester", status="completed", output=[ - MessageOutputItem( - id=f"msg_{i}", - content=[OutputTextContent(text=f"Response {i}")] - ) + MessageOutputItem(id=f"msg_{i}", content=[OutputTextContent(text=f"Response {i}")]) ], usage=ResponseUsage( - input_tokens=10, - output_tokens=5, - reasoning_tokens=0, - total_tokens=15 - ) + input_tokens=10, output_tokens=5, reasoning_tokens=0, total_tokens=15 + ), ) await history.add_response(conv_id, response) @@ -178,7 +161,7 @@ async def test_clear_conversation(): """Test clearing conversation history.""" history = ConversationHistory() - from src.responses.schemas import Response, ResponseUsage, MessageOutputItem, OutputTextContent + from src.responses.schemas import MessageOutputItem, OutputTextContent, Response, ResponseUsage conv_id = "test_clear" @@ -188,18 +171,8 @@ async def test_clear_conversation(): created_at=1234567890, model="lorem-tester", status="completed", - output=[ - MessageOutputItem( - id="msg_1", - content=[OutputTextContent(text="Test")] - ) - ], - usage=ResponseUsage( - input_tokens=10, - output_tokens=5, - reasoning_tokens=0, - total_tokens=15 - ) + output=[MessageOutputItem(id="msg_1", content=[OutputTextContent(text="Test")])], + usage=ResponseUsage(input_tokens=10, output_tokens=5, reasoning_tokens=0, total_tokens=15), ) await history.add_response(conv_id, response) @@ -227,7 +200,7 @@ async def test_conversation_count(): initial_count = await history.get_conversation_count() # Add responses to 3 different conversations - from src.responses.schemas import Response, ResponseUsage, MessageOutputItem, OutputTextContent + from src.responses.schemas import MessageOutputItem, OutputTextContent, Response, ResponseUsage for i in range(3): response = Response( @@ -235,18 +208,10 @@ async def test_conversation_count(): created_at=1234567890, model="lorem-tester", status="completed", - output=[ - MessageOutputItem( - id=f"msg_{i}", - content=[OutputTextContent(text="Test")] - ) - ], + output=[MessageOutputItem(id=f"msg_{i}", content=[OutputTextContent(text="Test")])], usage=ResponseUsage( - input_tokens=10, - output_tokens=5, - reasoning_tokens=0, - total_tokens=15 - ) + input_tokens=10, output_tokens=5, reasoning_tokens=0, total_tokens=15 + ), ) await history.add_response(f"conv_{i}", response) @@ -281,8 +246,8 @@ async def test_context_window_trimming(): # Create items that exceed limit items = [ "This is a long message " * 20, # ~480 chars = ~120 tokens - "Another message " * 10, # ~160 chars = ~40 tokens - "Short message" # ~13 chars = ~3 tokens + "Another message " * 10, # ~160 chars = ~40 tokens + "Short message", # ~13 chars = ~3 tokens ] # Trim with 10 token reserve @@ -355,14 +320,12 @@ async def test_service_conversation_helpers(): @pytest.mark.asyncio async def test_streaming_tracks_history(async_client): """Test that streaming responses also track conversation history.""" - from httpx import AsyncClient - import json request_data = { "model": "lorem-tester", "input": [{"role": "user", "content": "Hello"}], "metadata": {"conversation_id": "stream_test_001"}, - "stream": True + "stream": True, } async with async_client.stream( diff --git a/tests/responses/test_router.py b/tests/responses/test_router.py index 5ab522a..cddd0ba 100644 --- a/tests/responses/test_router.py +++ b/tests/responses/test_router.py @@ -3,6 +3,7 @@ Tests for Responses API router. """ import json + import pytest from fastapi.testclient import TestClient from httpx import AsyncClient @@ -13,10 +14,8 @@ def test_create_response_non_streaming(client: TestClient) -> None: """Test non-streaming response creation.""" request_data = { "model": "lorem-tester", - "input": [ - {"role": "user", "content": "Hello, how are you?"} - ], - "stream": False + "input": [{"role": "user", "content": "Hello, how are you?"}], + "stream": False, } response = client.post("/v1/responses", json=request_data) @@ -48,14 +47,9 @@ def test_create_response_with_reasoning(client: TestClient) -> None: """Test response with reasoning enabled.""" request_data = { "model": "lorem-tester", - "input": [ - {"role": "user", "content": "Explain something"} - ], - "reasoning": { - "effort": "medium", - "summary": "auto" - }, - "stream": False + "input": [{"role": "user", "content": "Explain something"}], + "reasoning": {"effort": "medium", "summary": "auto"}, + "stream": False, } response = client.post("/v1/responses", json=request_data) @@ -88,17 +82,11 @@ def test_create_response_with_tools(client: TestClient) -> None: """Test response with tools available.""" request_data = { "model": "lorem-tester", - "input": [ - {"role": "user", "content": "Use a tool to help"} - ], + "input": [{"role": "user", "content": "Use a tool to help"}], "tools": [ - { - "name": "search_knowledge", - "description": "Search knowledge base", - "parameters": {} - } + {"name": "search_knowledge", "description": "Search knowledge base", "parameters": {}} ], - "stream": False + "stream": False, } response = client.post("/v1/responses", json=request_data) @@ -117,10 +105,8 @@ def test_create_response_invalid_model(client: TestClient) -> None: """Test response with non-existent model.""" request_data = { "model": "nonexistent-model", - "input": [ - {"role": "user", "content": "Hello"} - ], - "stream": False + "input": [{"role": "user", "content": "Hello"}], + "stream": False, } response = client.post("/v1/responses", json=request_data) @@ -134,11 +120,7 @@ def test_create_response_invalid_model(client: TestClient) -> None: def test_create_response_validation_error(client: TestClient) -> None: """Test response with invalid request data.""" # Missing required 'model' field - request_data = { - "input": [ - {"role": "user", "content": "Hello"} - ] - } + request_data = {"input": [{"role": "user", "content": "Hello"}]} response = client.post("/v1/responses", json=request_data) @@ -151,10 +133,8 @@ async def test_create_response_streaming(async_client: AsyncClient) -> None: """Test streaming response creation.""" request_data = { "model": "lorem-tester", - "input": [ - {"role": "user", "content": "Hello"} - ], - "stream": True + "input": [{"role": "user", "content": "Hello"}], + "stream": True, } async with async_client.stream( @@ -200,14 +180,9 @@ async def test_create_response_streaming_with_reasoning(async_client: AsyncClien """Test streaming response with reasoning.""" request_data = { "model": "lorem-tester", - "input": [ - {"role": "user", "content": "Explain"} - ], - "reasoning": { - "effort": "medium", - "summary": "auto" - }, - "stream": True + "input": [{"role": "user", "content": "Explain"}], + "reasoning": {"effort": "medium", "summary": "auto"}, + "stream": True, } async with async_client.stream( @@ -238,16 +213,12 @@ async def test_create_response_streaming_with_reasoning(async_client: AsyncClien # Should have reasoning events reasoning_events = [ - e for e in events - if e["event"] == "response.reasoning_summary_text.delta" + e for e in events if e["event"] == "response.reasoning_summary_text.delta" ] assert len(reasoning_events) > 0 # Should have output text events - output_events = [ - e for e in events - if e["event"] == "response.output_text.delta" - ] + output_events = [e for e in events if e["event"] == "response.output_text.delta"] assert len(output_events) > 0 @@ -258,7 +229,7 @@ def test_create_response_temperature_parameter(client: TestClient) -> None: "model": "lorem-tester", "input": [{"role": "user", "content": "Test"}], "temperature": 0.5, - "stream": False + "stream": False, } response = client.post("/v1/responses", json=request_data) @@ -273,7 +244,7 @@ def test_create_response_max_tokens_parameter(client: TestClient) -> None: "model": "lorem-tester", "input": [{"role": "user", "content": "Test"}], "max_output_tokens": 100, - "stream": False + "stream": False, } response = client.post("/v1/responses", json=request_data) diff --git a/tests/responses/test_streaming_delegation.py b/tests/responses/test_streaming_delegation.py index 111536a..762f5f4 100644 --- a/tests/responses/test_streaming_delegation.py +++ b/tests/responses/test_streaming_delegation.py @@ -53,9 +53,7 @@ class TestStreamDirectDelegation: ) as mock_delegate: gen = coordinator._stream_direct_delegation( user_message="Search for Docker info", - recommendation=SimpleNamespace( - recommended_capabilities=["librarian"] - ), + recommendation=SimpleNamespace(recommended_capabilities=["librarian"]), tracker=tracker, conversation_id="conv_1", conversation_history=HISTORY, @@ -66,9 +64,9 @@ class TestStreamDirectDelegation: first = await gen.__anext__() assert isinstance(first, ReasoningSummaryDelta) assert first.delta.strip() != "" - assert mock_delegate.await_count == 0, ( - "start think message must stream before the expert runs" - ) + assert ( + mock_delegate.await_count == 0 + ), "start think message must stream before the expert runs" second = await gen.__anext__() assert isinstance(second, ReasoningSummaryDone) @@ -101,9 +99,7 @@ class TestStreamDirectDelegation: event async for event in coordinator._stream_direct_delegation( user_message="And what services does it list?", - recommendation=SimpleNamespace( - recommended_capabilities=["librarian"] - ), + recommendation=SimpleNamespace(recommended_capabilities=["librarian"]), tracker=tracker, conversation_id="conv_1", conversation_history=HISTORY, @@ -139,9 +135,7 @@ class TestStreamDirectDelegation: event async for event in coordinator._stream_direct_delegation( user_message="Search for Docker info", - recommendation=SimpleNamespace( - recommended_capabilities=["librarian"] - ), + recommendation=SimpleNamespace(recommended_capabilities=["librarian"]), tracker=tracker, conversation_id="conv_1", results=results, @@ -172,9 +166,7 @@ class TestServiceDelegationContext: ) as mock_delegate: results = await _direct_delegation_with_results( user_message="And what services does it list?", - recommendation=SimpleNamespace( - recommended_capabilities=["librarian"] - ), + recommendation=SimpleNamespace(recommended_capabilities=["librarian"]), tracker=tracker, conversation_id="conv_1", conversation_history=HISTORY, diff --git a/tests/test_main.py b/tests/test_main.py index 9e546ef..6eda9ca 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -8,6 +8,7 @@ Tests: - Router registration - Exception handlers """ + import pytest from fastapi import FastAPI from fastapi.testclient import TestClient @@ -56,8 +57,9 @@ def test_cors_middleware(): # Better check: verify CORS functionality by checking middleware stack # The middleware should be present (wrapped as Middleware) - middleware_types = [type(m.cls).__name__ if hasattr(m, 'cls') else type(m).__name__ - for m in app.user_middleware] + middleware_types = [ + type(m.cls).__name__ if hasattr(m, "cls") else type(m).__name__ for m in app.user_middleware + ] assert "CORSMiddleware" in middleware_types or len(app.user_middleware) > 0 @@ -67,10 +69,7 @@ def test_cors_allows_all_origins(client: TestClient): # Make OPTIONS request with Origin header response = client.options( "/v1/models", - headers={ - "Origin": "http://example.com", - "Access-Control-Request-Method": "GET" - } + headers={"Origin": "http://example.com", "Access-Control-Request-Method": "GET"}, ) # Should allow CORS @@ -81,10 +80,7 @@ def test_cors_allows_all_origins(client: TestClient): @pytest.mark.unit def test_cors_allows_credentials(client: TestClient): """Test that CORS allows credentials.""" - response = client.get( - "/health", - headers={"Origin": "http://example.com"} - ) + response = client.get("/health", headers={"Origin": "http://example.com"}) # Should include CORS headers assert response.status_code == 200 @@ -96,10 +92,11 @@ def test_cors_allows_credentials(client: TestClient): @pytest.mark.unit def test_exception_handlers_registered(): """Test that custom exception handlers are registered.""" - from src.main import app - from src.core.exceptions import AppException from fastapi.exceptions import RequestValidationError + from src.core.exceptions import AppException + from src.main import app + # App should have exception handlers assert len(app.exception_handlers) > 0 @@ -119,8 +116,8 @@ def test_model_not_found_exception_handler(client: TestClient): json={ "model": "non-existent-model-xyz", "input": [{"role": "user", "content": "test"}], - "stream": False - } + "stream": False, + }, ) assert response.status_code == 404 @@ -137,8 +134,8 @@ def test_app_exception_handler(client: TestClient): json={ "model": "lorem-tester", "input": [], # Empty input might cause issues - "temperature": 3.0 # Invalid temperature - } + "temperature": 3.0, # Invalid temperature + }, ) # Should return validation error @@ -149,8 +146,8 @@ def test_app_exception_handler(client: TestClient): @pytest.mark.asyncio async def test_lifespan_startup(): """Test lifespan startup event.""" + from src.main import app - from contextlib import asynccontextmanager # The lifespan context manager should be defined assert hasattr(app.router, "lifespan_context") @@ -238,7 +235,7 @@ def test_validation_error_handler(client: TestClient): json={ # Missing "model" field "messages": [{"role": "user", "content": "test"}] - } + }, ) assert response.status_code == 422