diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index bbff4b8a..efd2dc2d 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -71,6 +71,10 @@ jobs: python-version: '3.11' extras: ".[dev,google-adk]" tests: "tests/test_google_adk_integration.py" + - integration: mcp + python-version: '3.11' + extras: ".[dev,mcp]" + tests: "tests/test_mcp_integration.py" steps: - name: Checkout code diff --git a/cascadeflow/__init__.py b/cascadeflow/__init__.py index 2656995f..040a1da7 100644 --- a/cascadeflow/__init__.py +++ b/cascadeflow/__init__.py @@ -99,6 +99,8 @@ def __getattr__(self, name: str): # Agent & result "CascadeAgent": (".agent", "CascadeAgent"), "CascadeResult": (".schema.result", "CascadeResult"), + "KnowledgeCache": (".context", "KnowledgeCache"), + "KnowledgeSnapshot": (".context", "KnowledgeSnapshot"), "agent": (".agent", None), # Providers "BaseProvider": (".providers", "BaseProvider"), diff --git a/cascadeflow/agent.py b/cascadeflow/agent.py index c9f1778e..1f4f6ff9 100644 --- a/cascadeflow/agent.py +++ b/cascadeflow/agent.py @@ -68,6 +68,7 @@ from .profiles import UserProfile from .rules.decision import RuleDecision +from .context import KnowledgeCache, KnowledgeSnapshot, PreparedKnowledge, provider_cache_kwargs from .core.cascade import WholeResponseCascade # Phase 2C: Interface module imports @@ -211,6 +212,7 @@ def __init__( # 🆕 v2.9: Tool Execution # ======================================================================== tool_executor: Optional[Any] = None, # ToolExecutor instance or async callable + knowledge_cache: Optional[KnowledgeCache] = None, ): """ Initialize cascade agent with dual streaming managers and cost calculator. @@ -270,6 +272,9 @@ def __init__( # 🆕 v2.9: Tool execution support self._tool_executor = tool_executor + # Knowledge selection is request-scoped. The cache only reuses the stable + # compiled prefix; it never keeps an implicit "currently active" snapshot. + self.knowledge_cache = knowledge_cache or KnowledgeCache() # Setup logging if verbose: @@ -757,6 +762,15 @@ def _normalize_messages( ) -> tuple[str, Optional[list[dict[str, Any]]]]: if messages: normalized = normalize_messages(messages) + # ``system_prompt`` and knowledge both create a system message even + # when callers use the simple string API. Preserve that string as + # the current user turn instead of accidentally sending only context. + if query.strip() and not ( + normalized + and normalized[-1].get("role") == "user" + and normalized[-1].get("content", "").strip() == query.strip() + ): + normalized.append({"role": "user", "content": query}) query_text = messages_to_prompt(normalized) return query_text, normalized return query, None @@ -774,6 +788,34 @@ def _apply_system_prompt( return base return [{"role": "system", "content": system_prompt}, *base] + def _apply_knowledge( + self, + messages: Optional[list[dict[str, Any]]], + knowledge: Optional["str | KnowledgeSnapshot"], + provider_kwargs: dict[str, Any], + ) -> tuple[Optional[list[dict[str, Any]]], Optional[PreparedKnowledge], bool]: + """Prepend an immutable knowledge snapshot without carrying prior state. + + The rendered prefix is identical for every model in a cascade. Provider + cache intent remains internal until the concrete provider is selected. + """ + if knowledge is None: + return messages, None, False + + prepared, local_hit = self.knowledge_cache.prepare(knowledge) + base = list(messages or []) + if base and base[0].get("role") == "system": + existing = str(base[0].get("content") or "").strip() + content = prepared.system_prefix + if existing: + content = f"{content}\n\n{existing}" + base[0] = {**base[0], "content": content} + else: + base.insert(0, {"role": "system", "content": prepared.system_prefix}) + + provider_kwargs["_cascadeflow_prepared_knowledge"] = prepared + return base, prepared, local_hit + async def _execute_tool_calls_parallel( self, tool_calls: list[dict[str, Any]] ) -> list[dict[str, Any]]: @@ -863,6 +905,7 @@ async def run( tools: Optional[list[dict[str, Any]]] = None, tool_choice: Optional[str] = None, messages: Optional[list[dict[str, Any]]] = None, + knowledge: Optional["str | KnowledgeSnapshot"] = None, max_steps: int = 5, user_tier: Optional[str] = None, # 🔄 OPTIONAL: v0.1.x backwards compatibility workflow: Optional[str] = None, @@ -888,6 +931,7 @@ async def run( tools: List of tools in universal format tool_choice: Control tool calling behavior messages: Optional multi-turn messages (role/content) + knowledge: Immutable request-scoped knowledge snapshot (or string) user_tier: OPTIONAL - User tier for tier-based routing (v0.1.x compat) workflow: OPTIONAL - Workflow profile name (legacy) kpi_flags: OPTIONAL - KPI routing flags (risk/compliance, etc.) @@ -910,6 +954,9 @@ async def run( timing_breakdown = {} system_prompt = kwargs.pop("system_prompt", None) messages = self._apply_system_prompt(messages, system_prompt) + messages, prepared_knowledge, knowledge_local_hit = self._apply_knowledge( + messages, knowledge, kwargs + ) query_text, normalized_messages = self._normalize_messages(query, messages) tools = normalize_tools(tools) @@ -1286,6 +1333,11 @@ async def run( channel=channel, ) + if prepared_knowledge is not None: + cascade_result.metadata["knowledge"] = prepared_knowledge.metadata( + local_cache_hit=knowledge_local_hit + ) + # Record metrics with corrected cost values from CostCalculator self.telemetry.record( result=cascade_result, @@ -1314,6 +1366,7 @@ async def run_streaming( tools: Optional[list[dict[str, Any]]] = None, tool_choice: Optional[str] = None, messages: Optional[list[dict[str, Any]]] = None, + knowledge: Optional["str | KnowledgeSnapshot"] = None, user_tier: Optional[str] = None, workflow: Optional[str] = None, kpi_flags: Optional[dict[str, Any]] = None, @@ -1340,6 +1393,7 @@ async def run_streaming( tools: List of tools in universal format tool_choice: Control tool calling behavior messages: Optional multi-turn messages (role/content) + knowledge: Immutable request-scoped knowledge snapshot (or string) user_tier: OPTIONAL - User tier for tier-based routing (v0.1.x compat) workflow: OPTIONAL - Workflow profile name (legacy) kpi_flags: OPTIONAL - KPI routing flags (risk/compliance, etc.) @@ -1356,6 +1410,9 @@ async def run_streaming( timing_breakdown = {} system_prompt = kwargs.pop("system_prompt", None) messages = self._apply_system_prompt(messages, system_prompt) + messages, prepared_knowledge, knowledge_local_hit = self._apply_knowledge( + messages, knowledge, kwargs + ) query_text, normalized_messages = self._normalize_messages(query, messages) tools = normalize_tools(tools) @@ -1575,6 +1632,11 @@ async def run_streaming( channel=channel, ) + if prepared_knowledge is not None: + cascade_result.metadata["knowledge"] = prepared_knowledge.metadata( + local_cache_hit=knowledge_local_hit + ) + # Record metrics with corrected cost values from CostCalculator self.telemetry.record( result=cascade_result, @@ -1602,6 +1664,7 @@ async def stream_events( tools: Optional[list[dict[str, Any]]] = None, tool_choice: Optional[str] = None, messages: Optional[list[dict[str, Any]]] = None, + knowledge: Optional["str | KnowledgeSnapshot"] = None, user_tier: Optional[str] = None, workflow: Optional[str] = None, kpi_flags: Optional[dict[str, Any]] = None, @@ -1625,6 +1688,7 @@ async def stream_events( tools: List of tools in universal format tool_choice: Control tool calling behavior messages: Optional multi-turn messages (role/content) + knowledge: Immutable request-scoped knowledge snapshot (or string) user_tier: OPTIONAL - User tier for tier-based routing (v0.1.x compat) workflow: OPTIONAL - Workflow profile name (legacy) kpi_flags: OPTIONAL - KPI routing flags (risk/compliance, etc.) @@ -1640,6 +1704,9 @@ async def stream_events( # Detect complexity system_prompt = kwargs.pop("system_prompt", None) messages = self._apply_system_prompt(messages, system_prompt) + messages, _prepared_knowledge, _knowledge_local_hit = self._apply_knowledge( + messages, knowledge, kwargs + ) query_text, normalized_messages = self._normalize_messages(query, messages) tools = normalize_tools(tools) complexity_metadata = {} @@ -2149,6 +2216,8 @@ async def _execute_direct_with_timing( """Execute direct routing with detailed timing and tool support.""" best_model = available_models[-1] if available_models else self.models[-1] provider = self._get_provider(best_model) + prepared_knowledge = kwargs.pop("_cascadeflow_prepared_knowledge", None) + kwargs.update(provider_cache_kwargs(best_model.provider, prepared_knowledge)) reason = ( "Forced direct routing" if force_direct @@ -2288,6 +2357,8 @@ async def _stream_direct_with_timing( """Stream directly from best model with timing tracking and tool support.""" best_model = available_models[-1] if available_models else self.models[-1] provider = self._get_provider(best_model) + prepared_knowledge = kwargs.pop("_cascadeflow_prepared_knowledge", None) + kwargs.update(provider_cache_kwargs(best_model.provider, prepared_knowledge)) reason = ( "Forced direct routing" if force_direct diff --git a/cascadeflow/context/__init__.py b/cascadeflow/context/__init__.py new file mode 100644 index 00000000..a7f673a6 --- /dev/null +++ b/cascadeflow/context/__init__.py @@ -0,0 +1,17 @@ +"""Provider-neutral conversation and knowledge context helpers.""" + +from .knowledge import ( + KnowledgeCache, + KnowledgeInput, + KnowledgeSnapshot, + PreparedKnowledge, + provider_cache_kwargs, +) + +__all__ = [ + "KnowledgeCache", + "KnowledgeInput", + "KnowledgeSnapshot", + "PreparedKnowledge", + "provider_cache_kwargs", +] diff --git a/cascadeflow/context/knowledge.py b/cascadeflow/context/knowledge.py new file mode 100644 index 00000000..a591ffd0 --- /dev/null +++ b/cascadeflow/context/knowledge.py @@ -0,0 +1,205 @@ +"""Versioned, provider-neutral knowledge handoff. + +Provider prompt caches are an optimization, not a source of conversation state. +Every request is given the complete immutable knowledge snapshot it selected. A +stable rendering then lets providers reuse prefix caches when they support them, +without changing correctness for providers that do not. +""" + +from __future__ import annotations + +import hashlib +import re +import threading +from collections import OrderedDict +from dataclasses import dataclass +from typing import Any, Literal, Optional, Union + +CacheTTL = Literal["5m", "1h"] + + +def _digest(value: str) -> str: + return hashlib.sha256(value.encode("utf-8")).hexdigest() + + +def _safe_label(value: str) -> str: + """Keep prompt metadata deterministic and free of delimiter injection.""" + return re.sub(r"[^a-zA-Z0-9._:/-]+", "-", value.strip()).strip("-") or "knowledge" + + +@dataclass(frozen=True) +class KnowledgeSnapshot: + """An immutable knowledge version selected for one request. + + ``key`` names the logical knowledge set. ``version`` should change whenever + its content changes; when omitted it is derived from the content digest. + This prevents a caller from accidentally reusing stale provider-cache state. + """ + + content: str + key: str = "default" + version: Optional[str] = None + enable_provider_cache: bool = True + cache_ttl: CacheTTL = "5m" + + def __post_init__(self) -> None: + if not isinstance(self.content, str) or not self.content.strip(): + raise ValueError("KnowledgeSnapshot.content must be a non-empty string") + if not isinstance(self.key, str) or not self.key.strip(): + raise ValueError("KnowledgeSnapshot.key must be a non-empty string") + if self.cache_ttl not in ("5m", "1h"): + raise ValueError("KnowledgeSnapshot.cache_ttl must be '5m' or '1h'") + + @property + def content_digest(self) -> str: + return _digest(self.content) + + @property + def resolved_version(self) -> str: + return _safe_label(self.version or self.content_digest[:16]) + + @property + def identity(self) -> str: + return f"{_safe_label(self.key)}:{self.resolved_version}" + + +@dataclass(frozen=True) +class PreparedKnowledge: + """Stable prompt material plus provider-cache routing metadata.""" + + identity: str + content_digest: str + system_prefix: str + enable_provider_cache: bool + cache_ttl: CacheTTL + + @property + def prompt_cache_key(self) -> str: + # Bind native caches to both the caller's version and the actual content. + # This remains safe even after local LRU eviction or process restart. + material = f"{self.identity}:{self.content_digest}" + return f"cascadeflow:knowledge:{_digest(material)[:24]}" + + def metadata(self, *, local_cache_hit: bool) -> dict[str, Any]: + return { + "identity": self.identity, + "content_digest": self.content_digest, + "local_cache_hit": local_cache_hit, + "provider_cache_enabled": self.enable_provider_cache, + "cache_ttl": self.cache_ttl, + } + + +KnowledgeInput = Union[str, KnowledgeSnapshot] + + +class KnowledgeCache: + """Small LRU for compiled immutable snapshots. + + This cache deliberately has no global "active knowledge" pointer. Selection is + request-scoped, which makes concurrent requests and knowledge switches safe. + """ + + def __init__(self, max_entries: int = 128): + if max_entries <= 0: + raise ValueError("max_entries must be positive") + self.max_entries = max_entries + self._entries: OrderedDict[tuple[str, str, str, bool, CacheTTL], PreparedKnowledge] = ( + OrderedDict() + ) + self._version_digests: dict[tuple[str, str], str] = {} + self._lock = threading.RLock() + self._hits = 0 + self._misses = 0 + self._evictions = 0 + + def prepare(self, value: KnowledgeInput) -> tuple[PreparedKnowledge, bool]: + snapshot = value if isinstance(value, KnowledgeSnapshot) else KnowledgeSnapshot(value) + version_key = (_safe_label(snapshot.key), snapshot.resolved_version) + digest = snapshot.content_digest + cache_key = ( + version_key[0], + version_key[1], + digest, + snapshot.enable_provider_cache, + snapshot.cache_ttl, + ) + + with self._lock: + prior_digest = self._version_digests.get(version_key) + if prior_digest is not None and prior_digest != digest: + raise ValueError( + "Knowledge version collision: the same key/version was reused with " + "different content. Change KnowledgeSnapshot.version to prevent stale context." + ) + + prepared = self._entries.get(cache_key) + if prepared is not None: + self._entries.move_to_end(cache_key) + self._hits += 1 + return prepared, True + + identity = f"{version_key[0]}:{version_key[1]}" + prefix = ( + f'\n' + f"{snapshot.content.strip()}\n" + "" + ) + prepared = PreparedKnowledge( + identity=identity, + content_digest=digest, + system_prefix=prefix, + enable_provider_cache=snapshot.enable_provider_cache, + cache_ttl=snapshot.cache_ttl, + ) + self._entries[cache_key] = prepared + self._version_digests[version_key] = digest + self._misses += 1 + + if len(self._entries) > self.max_entries: + evicted_key, evicted = self._entries.popitem(last=False) + if not any(item.identity == evicted.identity for item in self._entries.values()): + self._version_digests.pop((evicted_key[0], evicted_key[1]), None) + self._evictions += 1 + + return prepared, False + + def clear(self) -> None: + with self._lock: + self._entries.clear() + self._version_digests.clear() + + def get_stats(self) -> dict[str, int]: + with self._lock: + return { + "size": len(self._entries), + "hits": self._hits, + "misses": self._misses, + "evictions": self._evictions, + } + + +def provider_cache_kwargs(provider: str, prepared: Optional[PreparedKnowledge]) -> dict[str, Any]: + """Translate a neutral cache intent into supported provider request fields. + + Unknown and self-hosted providers receive no special arguments. They still get + the exact knowledge prefix, so behavior remains correct without native caching. + """ + if prepared is None or not prepared.enable_provider_cache: + return {} + + provider = provider.lower() + if provider == "openai": + return { + "prompt_cache_key": prepared.prompt_cache_key, + # Consumed by OpenAIProvider and never forwarded as an API field. + # GPT-5.6+ needs the exact end of the stable prefix so it can place + # an explicit breakpoint before the changing query suffix. + "_cascadeflow_knowledge_cache_prefix": prepared.system_prefix, + } + if provider == "anthropic": + cache_control: dict[str, str] = {"type": "ephemeral"} + if prepared.cache_ttl == "1h": + cache_control["ttl"] = "1h" + return {"cache_control": cache_control} + return {} diff --git a/cascadeflow/core/cascade.py b/cascadeflow/core/cascade.py index 4144c2de..ad7bbae1 100644 --- a/cascadeflow/core/cascade.py +++ b/cascadeflow/core/cascade.py @@ -48,6 +48,7 @@ from enum import Enum from typing import Any, Optional +from ..context import PreparedKnowledge, provider_cache_kwargs from ..quality import AdaptiveThreshold, ComparativeValidator, QualityConfig, QualityValidator from ..schema.config import ModelConfig from ..utils.messages import get_last_user_message, messages_to_prompt, normalize_messages @@ -636,6 +637,8 @@ async def _execute_tool_path( 4. Accept or escalate to verifier 5. Calculate costs using CostCalculator (with input tokens!) """ + prepared_knowledge = kwargs.pop("_cascadeflow_prepared_knowledge", None) + # Timing breakdown timing = { "draft_latency_ms": 0.0, @@ -683,6 +686,7 @@ async def _execute_tool_path( tools=tools, tool_choice=tool_choice, messages=messages, + prepared_knowledge=prepared_knowledge, ) timing["draft_latency_ms"] = (time.time() - draft_start) * 1000 @@ -704,6 +708,7 @@ async def _execute_tool_path( tools=tools, tool_choice=tool_choice, messages=messages, + prepared_knowledge=prepared_knowledge, ) timing["verifier_latency_ms"] = (time.time() - verifier_start) * 1000 timing["total_latency_ms"] = (time.time() - overall_start) * 1000 @@ -767,6 +772,7 @@ async def _execute_tool_path( tools=tools, tool_choice=tool_choice, messages=messages, + prepared_knowledge=prepared_knowledge, ) ) else: @@ -881,6 +887,7 @@ async def _execute_tool_path( tools=tools, tool_choice=tool_choice, messages=messages, + prepared_knowledge=prepared_knowledge, ) else: verifier_result = await verifier_task @@ -1020,6 +1027,8 @@ async def _execute_text_path( FIXED: Now uses CostCalculator for accurate cost tracking. FIXED: Now passes query to CostCalculator for input token counting! """ + prepared_knowledge = kwargs.pop("_cascadeflow_prepared_knowledge", None) + # Timing breakdown timing = { "draft_latency_ms": 0.0, @@ -1035,7 +1044,12 @@ async def _execute_text_path( # === PHASE 1: Generate Draft === draft_start = time.time() draft_result = await self._call_drafter( - query, max_tokens, temperature, tools=None, messages=messages + query, + max_tokens, + temperature, + tools=None, + messages=messages, + prepared_knowledge=prepared_knowledge, ) timing["draft_latency_ms"] = (time.time() - draft_start) * 1000 @@ -1049,7 +1063,12 @@ async def _execute_text_path( verifier_start = time.time() verifier_result = await self._call_verifier( - query, max_tokens, temperature, tools=None, messages=messages + query, + max_tokens, + temperature, + tools=None, + messages=messages, + prepared_knowledge=prepared_knowledge, ) timing["verifier_latency_ms"] = (time.time() - verifier_start) * 1000 timing["total_latency_ms"] = (time.time() - overall_start) * 1000 @@ -1102,7 +1121,14 @@ async def _execute_text_path( f"Draft confidence {raw_draft_confidence:.2f} < 0.75, starting verifier" ) verifier_task = asyncio.create_task( - self._call_verifier(query, max_tokens, temperature, tools=None, messages=messages) + self._call_verifier( + query, + max_tokens, + temperature, + tools=None, + messages=messages, + prepared_knowledge=prepared_knowledge, + ) ) else: verifier_task = None @@ -1247,7 +1273,12 @@ async def _execute_text_path( verifier_start = time.time() if verifier_task is None: verifier_result = await self._call_verifier( - query, max_tokens, temperature, tools=None, messages=messages + query, + max_tokens, + temperature, + tools=None, + messages=messages, + prepared_knowledge=prepared_knowledge, ) else: verifier_result = await verifier_task @@ -1314,6 +1345,7 @@ async def _call_drafter( tools: Optional[list[dict[str, Any]]] = None, tool_choice: Optional[str] = None, messages: Optional[list[dict[str, Any]]] = None, + prepared_knowledge: Optional[PreparedKnowledge] = None, ) -> Optional[dict[str, Any]]: """ Call drafter model with optional tool support. @@ -1321,6 +1353,7 @@ async def _call_drafter( """ try: provider = self._get_provider(self.drafter) + cache_kwargs = provider_cache_kwargs(self.drafter.provider, prepared_knowledge) # === CRITICAL FIX: Route to correct method based on tools === if tools: @@ -1334,6 +1367,7 @@ async def _call_drafter( model=self.drafter.name, max_tokens=max_tokens, temperature=temperature, + **cache_kwargs, ) else: # TEXT PATH: Use complete() with prompt format @@ -1345,6 +1379,7 @@ async def _call_drafter( temperature=temperature, logprobs=True, top_logprobs=5, + **cache_kwargs, ) return _convert_to_dict(result) @@ -1361,6 +1396,7 @@ async def _call_verifier( tools: Optional[list[dict[str, Any]]] = None, tool_choice: Optional[str] = None, messages: Optional[list[dict[str, Any]]] = None, + prepared_knowledge: Optional[PreparedKnowledge] = None, ) -> dict[str, Any]: """ Call verifier model with optional tool support. @@ -1368,6 +1404,7 @@ async def _call_verifier( """ try: provider = self._get_provider(self.verifier) + cache_kwargs = provider_cache_kwargs(self.verifier.provider, prepared_knowledge) # === CRITICAL FIX: Route to correct method based on tools === if tools: @@ -1381,6 +1418,7 @@ async def _call_verifier( model=self.verifier.name, max_tokens=max_tokens, temperature=temperature, + **cache_kwargs, ) else: # TEXT PATH: Use complete() with prompt format @@ -1392,6 +1430,7 @@ async def _call_verifier( temperature=temperature, logprobs=True, top_logprobs=5, + **cache_kwargs, ) return _convert_to_dict(result) diff --git a/cascadeflow/integrations/__init__.py b/cascadeflow/integrations/__init__.py index ed2cf58e..1e0272f4 100644 --- a/cascadeflow/integrations/__init__.py +++ b/cascadeflow/integrations/__init__.py @@ -19,6 +19,7 @@ from __future__ import annotations +import importlib.util from typing import TYPE_CHECKING @@ -334,6 +335,14 @@ def __repr__(self): extract_cascadeflow_skill_metadata = _hermes_missing profile_from_skill_metadata = _hermes_missing +# ═══════════════════════════════════════════════════ +# Model Context Protocol +# ═══════════════════════════════════════════════════ + +from .mcp import KnowledgeResolver, create_mcp_server + +MCP_AVAILABLE = importlib.util.find_spec("mcp") is not None + # ═══════════════════════════════════════════════════ # Exports & Capabilities @@ -483,6 +492,9 @@ def __repr__(self): ] ) +if MCP_AVAILABLE: + __all__.extend(["MCP_AVAILABLE", "KnowledgeResolver", "create_mcp_server"]) + # Integration capabilities INTEGRATION_CAPABILITIES = { "litellm": LITELLM_AVAILABLE, @@ -495,6 +507,7 @@ def __repr__(self): "google_adk": GOOGLE_ADK_AVAILABLE, "pydantic_ai": PYDANTIC_AI_AVAILABLE, "hermes": HERMES_AVAILABLE, + "mcp": MCP_AVAILABLE, } @@ -523,4 +536,5 @@ def get_integration_info(): "google_adk_available": GOOGLE_ADK_AVAILABLE, "pydantic_ai_available": PYDANTIC_AI_AVAILABLE, "hermes_available": HERMES_AVAILABLE, + "mcp_available": MCP_AVAILABLE, } diff --git a/cascadeflow/integrations/mcp.py b/cascadeflow/integrations/mcp.py new file mode 100644 index 00000000..e8cfb635 --- /dev/null +++ b/cascadeflow/integrations/mcp.py @@ -0,0 +1,133 @@ +"""MCP tool adapter for running cascadeflow from compatible chat hosts. + +The adapter intentionally keeps knowledge resolution server-side. MCP clients send +only a stable knowledge identifier and, when necessary, a concise conversation +handoff. This avoids duplicating an entire knowledge base in every tool call. +""" + +from __future__ import annotations + +import inspect +from typing import Any, Awaitable, Callable, Optional, Union + +from ..context import KnowledgeInput, KnowledgeSnapshot + +KnowledgeResolverResult = Union[KnowledgeInput, Awaitable[KnowledgeInput]] +KnowledgeResolver = Callable[[str, Optional[str]], KnowledgeResolverResult] + + +def _load_fastmcp() -> Any: + try: + from mcp.server.fastmcp import FastMCP + except ImportError as exc: # pragma: no cover - exercised without the optional extra + raise ImportError( + "The MCP integration requires Python 3.10+ and the MCP SDK. " + "Install it with: pip install 'cascadeflow[mcp]'" + ) from exc + return FastMCP + + +async def _resolve_knowledge( + resolver: KnowledgeResolver, + key: str, + version: Optional[str], +) -> KnowledgeInput: + resolved = resolver(key, version) + if inspect.isawaitable(resolved): + resolved = await resolved + if isinstance(resolved, str): + return KnowledgeSnapshot(content=resolved, key=key, version=version) + if not isinstance(resolved, KnowledgeSnapshot): + raise TypeError("knowledge_resolver must return str or KnowledgeSnapshot") + return resolved + + +def create_mcp_server( + agent: Any, + *, + knowledge_resolver: Optional[KnowledgeResolver] = None, + name: str = "cascadeflow", + max_context_chars: int = 12_000, +) -> Any: + """Create a tool-first MCP server backed by a configured ``CascadeAgent``. + + The returned FastMCP instance can run over stdio for local desktop clients or + Streamable HTTP for remote hosts. Knowledge content stays behind the server; + clients select it with ``knowledge_key`` and optional ``knowledge_version``. + """ + if max_context_chars <= 0: + raise ValueError("max_context_chars must be positive") + + FastMCP = _load_fastmcp() + server = FastMCP( + name, + instructions=( + "Use cascadeflow_run when the user wants a cost-aware answer from " + "cascadeflow. Select server-side knowledge with knowledge_key and send " + "conversation_context only as a concise factual handoff when prior turns " + "are required." + ), + stateless_http=True, + json_response=True, + ) + + @server.tool( + title="Run cascadeflow", + annotations={ + "readOnlyHint": True, + "destructiveHint": False, + "openWorldHint": True, + }, + ) + async def cascadeflow_run( + query: str, + knowledge_key: Optional[str] = None, + knowledge_version: Optional[str] = None, + conversation_context: Optional[str] = None, + ) -> dict[str, Any]: + """Run a query through cascadeflow's cost-aware model routing. + + Use knowledge_key to select server-side knowledge. Send + conversation_context only when the query depends on prior turns, and keep + it to a concise factual handoff rather than the complete chat transcript. + """ + if not query.strip(): + raise ValueError("query must be non-empty") + if conversation_context and len(conversation_context) > max_context_chars: + raise ValueError( + f"conversation_context exceeds {max_context_chars} characters; " + "send a concise relevant handoff" + ) + + knowledge: Optional[KnowledgeInput] = None + if knowledge_key: + if knowledge_resolver is None: + raise ValueError("knowledge_key requires a configured knowledge_resolver") + knowledge = await _resolve_knowledge( + knowledge_resolver, knowledge_key, knowledge_version + ) + + messages = None + if conversation_context and conversation_context.strip(): + messages = [ + {"role": "user", "content": conversation_context.strip()}, + {"role": "user", "content": query}, + ] + + result = await agent.run(query, messages=messages, knowledge=knowledge) + return { + "content": result.content, + "model_used": result.model_used, + "cascaded": result.cascaded, + "draft_accepted": result.draft_accepted, + "routing_strategy": result.routing_strategy, + "total_cost": result.total_cost, + "cost_saved": result.cost_saved, + "latency_ms": result.latency_ms, + "knowledge": result.metadata.get("knowledge"), + } + + return server + + +__all__ = ["KnowledgeResolver", "create_mcp_server"] diff --git a/cascadeflow/providers/anthropic.py b/cascadeflow/providers/anthropic.py index 0b72eaa8..e95cf127 100644 --- a/cascadeflow/providers/anthropic.py +++ b/cascadeflow/providers/anthropic.py @@ -11,6 +11,80 @@ from ..exceptions import ModelError, ProviderError from .base import BaseProvider, HttpConfig, ModelResponse, RetryConfig +_ANTHROPIC_PRICING: dict[str, dict[str, float]] = { + # USD per million tokens: input / output. + "claude-fable-5": {"input": 10.0, "output": 50.0}, + "claude-mythos-5": {"input": 10.0, "output": 50.0}, + "claude-opus-5": {"input": 5.0, "output": 25.0}, + "claude-opus-4-8": {"input": 5.0, "output": 25.0}, + "claude-opus-4-7": {"input": 5.0, "output": 25.0}, + "claude-opus-4-6": {"input": 5.0, "output": 25.0}, + "claude-opus-4-5": {"input": 5.0, "output": 25.0}, + "claude-opus-4.1": {"input": 15.0, "output": 75.0}, + "claude-opus-4": {"input": 15.0, "output": 75.0}, + # Introductory price through 2026-08-31; invoices remain authoritative. + "claude-sonnet-5": {"input": 2.0, "output": 10.0}, + "claude-sonnet-4-6": {"input": 3.0, "output": 15.0}, + "claude-sonnet-4-5": {"input": 3.0, "output": 15.0}, + "claude-sonnet-4.5": {"input": 3.0, "output": 15.0}, + "claude-sonnet-4": {"input": 3.0, "output": 15.0}, + "claude-haiku-4-5": {"input": 1.0, "output": 5.0}, + "claude-haiku-4.5": {"input": 1.0, "output": 5.0}, + "claude-3-7-sonnet": {"input": 3.0, "output": 15.0}, + "claude-3-5-sonnet": {"input": 3.0, "output": 15.0}, + "claude-sonnet-3-5": {"input": 3.0, "output": 15.0}, + "claude-3-5-haiku": {"input": 0.8, "output": 4.0}, + "claude-haiku-3-5": {"input": 0.8, "output": 4.0}, + "claude-3-opus": {"input": 15.0, "output": 75.0}, + "claude-3-sonnet": {"input": 3.0, "output": 15.0}, + "claude-3-haiku": {"input": 0.25, "output": 1.25}, +} + + +def _anthropic_model_pricing(model: str) -> dict[str, float]: + model_lower = model.lower().rsplit("/", 1)[-1] + matches = [prefix for prefix in _ANTHROPIC_PRICING if model_lower.startswith(prefix)] + if not matches: + return {"input": 3.0, "output": 15.0} + return _ANTHROPIC_PRICING[max(matches, key=len)] + + +def _cache_usage(usage: dict[str, Any]) -> dict[str, int]: + creation = usage.get("cache_creation") or {} + if not isinstance(creation, dict): + creation = {} + return { + "cached_input_tokens": int(usage.get("cache_read_input_tokens") or 0), + "cache_write_input_tokens": int(usage.get("cache_creation_input_tokens") or 0), + "cache_write_5m_input_tokens": int(creation.get("ephemeral_5m_input_tokens") or 0), + "cache_write_1h_input_tokens": int(creation.get("ephemeral_1h_input_tokens") or 0), + } + + +def _anthropic_usage_cost( + *, model: str, prompt_tokens: int, completion_tokens: int, usage: dict[str, Any] +) -> float: + """Calculate Anthropic cost from its disjoint input/cache usage categories.""" + pricing = _anthropic_model_pricing(model) + cache = _cache_usage(usage) + read_tokens = cache["cached_input_tokens"] + write_tokens = cache["cache_write_input_tokens"] + write_5m = min(cache["cache_write_5m_input_tokens"], write_tokens) + write_1h = min(cache["cache_write_1h_input_tokens"], max(write_tokens - write_5m, 0)) + unclassified_writes = max(write_tokens - write_5m - write_1h, 0) + + input_units = ( + max(prompt_tokens, 0) + + (read_tokens * 0.1) + + (write_5m * 1.25) + + (write_1h * 2.0) + + (unclassified_writes * 1.25) + ) + return ( + input_units * pricing["input"] + max(completion_tokens, 0) * pricing["output"] + ) / 1_000_000 + + # ============================================================================== # REASONING MODEL SUPPORT # ============================================================================== @@ -266,6 +340,30 @@ def _strip_internal_kwargs(self, extra: dict[str, Any]) -> dict[str, Any]: extra.pop(k, None) return extra + def _calculate_response_cost( + self, + *, + model: str, + prompt_tokens: int, + completion_tokens: int, + usage: dict[str, Any], + ) -> float: + """Use Anthropic cache categories when present, preserving LiteLLM otherwise.""" + cache = _cache_usage(usage) + if cache["cached_input_tokens"] or cache["cache_write_input_tokens"]: + return _anthropic_usage_cost( + model=model, + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + usage=usage, + ) + return self.calculate_accurate_cost( + model=model, + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + total_tokens=prompt_tokens + completion_tokens, + ) + def _convert_tools_to_anthropic(self, tools: list[dict[str, Any]]) -> list[dict[str, Any]]: """ Convert tools from universal format to Anthropic format. @@ -620,7 +718,7 @@ async def complete_with_tools( """ start_time = time.time() - self._strip_internal_kwargs(dict(kwargs)) + kwargs = self._strip_internal_kwargs(dict(kwargs)) # Convert tools to Anthropic format anthropic_tools = self._convert_tools_to_anthropic(tools) if tools else None @@ -687,11 +785,11 @@ async def complete_with_tools( latency_ms = (time.time() - start_time) * 1000 # Calculate cost using LiteLLM if available, otherwise fallback - cost = self.calculate_accurate_cost( + cost = self._calculate_response_cost( model=model, prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, - total_tokens=tokens_used, + usage=usage, ) # ============================================================ @@ -793,6 +891,7 @@ async def complete_with_tools( "id": data.get("id"), "prompt_tokens": prompt_tokens, "completion_tokens": completion_tokens, + **_cache_usage(usage), "has_tool_calls": bool(tool_calls), # NEW: Add confidence analysis details "query": user_query, @@ -963,11 +1062,11 @@ async def _complete_impl( latency_ms = (time.time() - start_time) * 1000 # Calculate cost using LiteLLM if available, otherwise fallback - cost = self.calculate_accurate_cost( + cost = self._calculate_response_cost( model=model, prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, - total_tokens=tokens_used, + usage=usage, ) # ============================================================ @@ -1019,6 +1118,7 @@ async def _complete_impl( "id": data.get("id"), "prompt_tokens": prompt_tokens, "completion_tokens": completion_tokens, + **_cache_usage(usage), # NEW: Add confidence analysis details for test validation "query": prompt, "confidence_method": confidence_method, @@ -1254,40 +1354,9 @@ def estimate_cost(self, tokens: int, model: str) -> float: Returns: Estimated cost in USD (blended average) """ - # Anthropic pricing per million tokens (October 2025) - # Format: Blended rate (50% input, 50% output) - # Official rates: Input / Output per MTok - rates = { - # Claude 4.x Series - "claude-opus-4-6": 15.0, # $5 in + $25 out = $15 blended - "claude-opus-4-5": 15.0, # $5 in + $25 out = $15 blended - "claude-opus-4.1": 45.0, # $15 in + $75 out = $45 blended - "claude-opus-4": 45.0, # $15 in + $75 out = $45 blended - "claude-sonnet-4-5": 9.0, # $3 in + $15 out = $9 blended - "claude-sonnet-4.5": 9.0, # $3 in + $15 out = $9 blended - "claude-sonnet-4": 9.0, # $3 in + $15 out = $9 blended - "claude-haiku-4-5": 3.0, # $1 in + $5 out = $3 blended - "claude-haiku-4.5": 3.0, # $1 in + $5 out = $3 blended - # Claude 3.5 Series - "claude-3-5-sonnet": 9.0, # $3 in + $15 out = $9 blended - "claude-sonnet-3-5": 9.0, # $3 in + $15 out = $9 blended (alternative naming) - "claude-3-5-haiku": 3.0, # $1 in + $5 out = $3 blended - "claude-haiku-3-5": 3.0, # $1 in + $5 out = $3 blended (alternative naming) - # Claude 3 Series - "claude-3-opus": 45.0, # $15 in + $75 out = $45 blended - "claude-3-sonnet": 9.0, # $3 in + $15 out = $9 blended - "claude-3-haiku": 0.75, # $0.25 in + $1.25 out = $0.75 blended - } - - model_lower = model.lower() - - # Find matching rate (try exact match first, then prefix) - for model_prefix, rate in rates.items(): - if model_lower.startswith(model_prefix): - return (tokens / 1_000_000) * rate - - # Default to Sonnet pricing if unknown (most common model) - return (tokens / 1_000_000) * 9.0 + pricing = _anthropic_model_pricing(model) + blended_rate = (pricing["input"] + pricing["output"]) / 2 + return (tokens / 1_000_000) * blended_rate async def __aenter__(self): """Async context manager entry.""" diff --git a/cascadeflow/providers/openai.py b/cascadeflow/providers/openai.py index 968cbe35..14b88ce1 100644 --- a/cascadeflow/providers/openai.py +++ b/cascadeflow/providers/openai.py @@ -2,6 +2,7 @@ import json import os +import re import time from collections.abc import AsyncIterator from typing import Any, Optional @@ -11,6 +12,101 @@ from ..exceptions import ModelError, ProviderError from .base import BaseProvider, HttpConfig, ModelResponse, RetryConfig +# Standard API pricing in USD per 1K tokens. GPT-5.6 rates are from the +# 2026-08 OpenAI pricebook; older entries preserve CascadeFlow's existing rates. +_OPENAI_PRICING: dict[str, dict[str, float]] = { + "gpt-5.6-sol": {"input": 0.005, "output": 0.030}, + "gpt-5.6-terra": {"input": 0.002, "output": 0.012}, + "gpt-5.6-luna": {"input": 0.0002, "output": 0.0012}, + "gpt-5.6": {"input": 0.005, "output": 0.030}, + "gpt-5-chat-latest": {"input": 0.00125, "output": 0.010}, + "gpt-5-mini": {"input": 0.00025, "output": 0.002}, + "gpt-5-nano": {"input": 0.00005, "output": 0.0004}, + "gpt-5": {"input": 0.00125, "output": 0.010}, + "gpt-4o-mini": {"input": 0.00015, "output": 0.0006}, + "gpt-4o": {"input": 0.0025, "output": 0.010}, + "o1-2024-12-17": {"input": 0.015, "output": 0.060}, + "o1-preview": {"input": 0.015, "output": 0.060}, + "o1-mini": {"input": 0.003, "output": 0.012}, + "o1": {"input": 0.015, "output": 0.060}, + "o3-mini": {"input": 0.001, "output": 0.005}, + "gpt-4-turbo": {"input": 0.010, "output": 0.030}, + "gpt-4": {"input": 0.030, "output": 0.060}, + "gpt-3.5-turbo": {"input": 0.0005, "output": 0.0015}, +} + + +def _openai_model_pricing(model: str) -> dict[str, float]: + """Return the longest-prefix price match for a versioned model name.""" + model_lower = model.lower().rsplit("/", 1)[-1] + matches = [prefix for prefix in _OPENAI_PRICING if model_lower.startswith(prefix)] + if not matches: + return {"input": 0.030, "output": 0.060} + return _OPENAI_PRICING[max(matches, key=len)] + + +def _cached_input_tokens(usage: dict[str, Any]) -> int: + details = usage.get("input_tokens_details") or usage.get("prompt_tokens_details") or {} + if not isinstance(details, dict): + return 0 + return int(details.get("cached_tokens") or 0) + + +def _cache_write_input_tokens(usage: dict[str, Any]) -> int: + """Return billable prompt-cache writes across current OpenAI response shapes.""" + details = usage.get("input_tokens_details") or usage.get("prompt_tokens_details") or {} + detail_value = details.get("cache_write_tokens") if isinstance(details, dict) else 0 + return int(usage.get("cache_write_tokens") or detail_value or 0) + + +def _uses_explicit_prompt_cache_breakpoints(model: str) -> bool: + """GPT-5.6+ rejects the older implicit-prefix cost assumptions.""" + model_name = (model or "").lower().rsplit("/", 1)[-1] + match = re.match(r"^gpt-(\d+)(?:\.(\d+))?", model_name) + if not match: + return False + major = int(match.group(1)) + minor = int(match.group(2) or 0) + return (major, minor) >= (5, 6) + + +def _cache_content_blocks(text: str, stable_prefix: str, *, block_type: str) -> Optional[list]: + """Split text immediately after knowledge and mark that stable prefix.""" + start = text.find(stable_prefix) + if start < 0: + return None + split_at = start + len(stable_prefix) + blocks: list[dict[str, Any]] = [ + { + "type": block_type, + "text": text[:split_at], + "prompt_cache_breakpoint": {"mode": "explicit"}, + } + ] + if text[split_at:]: + blocks.append({"type": block_type, "text": text[split_at:]}) + return blocks + + +def _openai_usage_cost( + *, model: str, prompt_tokens: int, completion_tokens: int, usage: dict[str, Any] +) -> float: + """Calculate cache-aware OpenAI cost from provider-reported token categories.""" + pricing = _openai_model_pricing(model) + cached_tokens = min(_cached_input_tokens(usage), max(prompt_tokens, 0)) + write_tokens = min(_cache_write_input_tokens(usage), max(prompt_tokens - cached_tokens, 0)) + uncached_tokens = max(prompt_tokens - cached_tokens - write_tokens, 0) + write_multiplier = 1.25 if _uses_explicit_prompt_cache_breakpoints(model) else 1.0 + + input_cost = ( + (uncached_tokens + (cached_tokens * 0.1) + (write_tokens * write_multiplier)) + * pricing["input"] + / 1000 + ) + output_cost = max(completion_tokens, 0) * pricing["output"] / 1000 + return input_cost + output_cost + + # ============================================================================== # REASONING MODEL SUPPORT # ============================================================================== @@ -548,6 +644,77 @@ def _convert_messages_to_responses_input( instructions = "\n".join(instructions_parts).strip() if instructions_parts else "" return input_messages, instructions or None + def _apply_responses_cache_breakpoint( + self, + input_messages: list[dict[str, Any]], + instructions: Optional[str], + stable_prefix: Optional[str], + ) -> tuple[list[dict[str, Any]], Optional[str]]: + """Mark the stable knowledge prefix for GPT-5.6+ Responses requests.""" + if not stable_prefix: + return input_messages, instructions + + if instructions: + blocks = _cache_content_blocks(instructions, stable_prefix, block_type="input_text") + if blocks: + return [ + {"role": "developer", "content": blocks}, + *input_messages, + ], None + + for index, message in enumerate(input_messages): + content = message.get("content") + if not isinstance(content, str): + continue + blocks = _cache_content_blocks(content, stable_prefix, block_type="input_text") + if blocks: + updated = list(input_messages) + updated[index] = {**message, "content": blocks} + return updated, instructions + return input_messages, instructions + + def _calculate_response_cost( + self, + *, + model: str, + prompt_tokens: int, + completion_tokens: int, + usage: dict[str, Any], + ) -> float: + """Use cache categories when present, preserving LiteLLM otherwise.""" + if _cached_input_tokens(usage) or _cache_write_input_tokens(usage): + return _openai_usage_cost( + model=model, + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + usage=usage, + ) + return self.calculate_accurate_cost( + model=model, + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + total_tokens=prompt_tokens + completion_tokens, + ) + + def _apply_chat_cache_breakpoint( + self, + messages: list[dict[str, Any]], + stable_prefix: Optional[str], + ) -> list[dict[str, Any]]: + """Mark the stable knowledge prefix for GPT-5.6+ Chat Completions.""" + if not stable_prefix: + return messages + for index, message in enumerate(messages): + content = message.get("content") + if not isinstance(content, str): + continue + blocks = _cache_content_blocks(content, stable_prefix, block_type="text") + if blocks: + updated = list(messages) + updated[index] = {**message, "content": blocks} + return updated + return messages + def _parse_responses_output( self, data: dict[str, Any] ) -> tuple[str, Optional[list[dict[str, Any]]], str]: @@ -735,6 +902,12 @@ async def complete_with_tools( extra = dict(kwargs) extra.pop("max_tokens", None) extra.pop("max_completion_tokens", None) + stable_cache_prefix = extra.pop("_cascadeflow_knowledge_cache_prefix", None) + explicit_cache = bool(stable_cache_prefix) and _uses_explicit_prompt_cache_breakpoints( + model + ) + if explicit_cache: + extra["prompt_cache_options"] = {"mode": "explicit"} extra_tool_choice = extra.pop("tool_choice", None) if tool_choice is None: tool_choice = extra_tool_choice @@ -742,6 +915,10 @@ async def complete_with_tools( # Prefer Responses API for GPT-5 (and optionally via env override). if self._should_use_responses_api(model): input_messages, instructions = self._convert_messages_to_responses_input(messages) + if explicit_cache: + input_messages, instructions = self._apply_responses_cache_breakpoint( + input_messages, instructions, stable_cache_prefix + ) max_out = self._effective_responses_max_output_tokens(model, max_tokens) payload: dict[str, Any] = { @@ -793,11 +970,11 @@ async def complete_with_tools( latency_ms = (time.time() - start_time) * 1000 - cost = self.calculate_accurate_cost( + cost = self._calculate_response_cost( model=model, prompt_tokens=int(prompt_tokens), completion_tokens=int(completion_tokens), - total_tokens=int(tokens_used), + usage=usage, ) if tool_calls: @@ -819,6 +996,8 @@ async def complete_with_tools( "finish_reason": finish_reason, "prompt_tokens": int(prompt_tokens), "completion_tokens": int(completion_tokens), + "cached_input_tokens": _cached_input_tokens(usage), + "cache_write_input_tokens": _cache_write_input_tokens(usage), "has_tool_calls": bool(tool_calls), "confidence_method": confidence_method, "tool_choice_reasoning": ( @@ -867,12 +1046,8 @@ async def complete_with_tools( # Build request payload with reasoning-model compatibility model_info = get_reasoning_model_info(model) is_gpt5 = model.lower().startswith("gpt-5") - extra = dict(kwargs) - extra.pop("max_tokens", None) - extra.pop("max_completion_tokens", None) - extra_tool_choice = extra.pop("tool_choice", None) - if tool_choice is None: - tool_choice = extra_tool_choice + if explicit_cache: + messages = self._apply_chat_cache_breakpoint(messages, stable_cache_prefix) payload = { "model": model, @@ -921,11 +1096,11 @@ async def complete_with_tools( latency_ms = (time.time() - start_time) * 1000 # Calculate cost (automatically uses LiteLLM if available) - cost = self.calculate_accurate_cost( + cost = self._calculate_response_cost( model=model, prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, - total_tokens=tokens_used, + usage=data.get("usage") or {}, ) # Parse tool calls if present @@ -961,6 +1136,8 @@ async def complete_with_tools( "finish_reason": choice["finish_reason"], "prompt_tokens": prompt_tokens, "completion_tokens": completion_tokens, + "cached_input_tokens": _cached_input_tokens(data.get("usage") or {}), + "cache_write_input_tokens": _cache_write_input_tokens(data.get("usage") or {}), "has_tool_calls": bool(tool_calls), "confidence_method": confidence_method, # ← NEW! "tool_choice_reasoning": ( @@ -1052,6 +1229,14 @@ async def _complete_impl( # Extract logprobs parameters logprobs_enabled = kwargs.pop("logprobs", False) top_logprobs = kwargs.pop("top_logprobs", 5) # Default to 5 + stable_cache_prefix = kwargs.pop("_cascadeflow_knowledge_cache_prefix", None) + explicit_cache = bool(stable_cache_prefix) and _uses_explicit_prompt_cache_breakpoints( + model + ) + if explicit_cache: + # Disable the implicit latest-message write. Otherwise each changing + # query can create a new billable GPT-5.6 cache entry. + kwargs["prompt_cache_options"] = {"mode": "explicit"} # Get reasoning model info for auto-configuration model_info = get_reasoning_model_info(model) @@ -1065,6 +1250,9 @@ async def _complete_impl( # Prepend system prompt to first user message prompt = f"{system_prompt}\n\n{prompt}" messages.append({"role": "user", "content": prompt}) + uses_responses_api = self._should_use_responses_api(model) + if explicit_cache and not uses_responses_api: + messages = self._apply_chat_cache_breakpoint(messages, stable_cache_prefix) # Check if this is GPT-5 model for correct token parameter is_gpt5 = model.lower().startswith("gpt-5") @@ -1099,8 +1287,12 @@ async def _complete_impl( payload["top_logprobs"] = min(top_logprobs, 20) # OpenAI max is 20 try: - if self._should_use_responses_api(model): + if uses_responses_api: input_messages, instructions = self._convert_messages_to_responses_input(messages) + if explicit_cache: + input_messages, instructions = self._apply_responses_cache_breakpoint( + input_messages, instructions, stable_cache_prefix + ) max_out = self._effective_responses_max_output_tokens(model, max_tokens) responses_payload: dict[str, Any] = { "model": model, @@ -1111,6 +1303,13 @@ async def _complete_impl( responses_payload["instructions"] = instructions if model_info.supports_temperature: responses_payload["temperature"] = temperature + for cache_field in ( + "prompt_cache_key", + "prompt_cache_retention", + "prompt_cache_options", + ): + if cache_field in payload: + responses_payload[cache_field] = payload[cache_field] response = await self.client.post( f"{self.base_url}/responses", json=responses_payload @@ -1145,11 +1344,11 @@ async def _complete_impl( latency_ms = (time.time() - start_time) * 1000 - cost = self.estimate_cost( - tokens_used, - model, + cost = self._calculate_response_cost( + model=model, prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, + usage=usage, ) metadata_for_confidence = { @@ -1182,6 +1381,8 @@ async def _complete_impl( "finish_reason": finish_reason, "prompt_tokens": prompt_tokens, "completion_tokens": completion_tokens, + "cached_input_tokens": _cached_input_tokens(usage), + "cache_write_input_tokens": _cache_write_input_tokens(usage), "query": prompt, "confidence_method": confidence_method, "confidence_components": confidence_components, @@ -1224,8 +1425,11 @@ async def _complete_impl( latency_ms = (time.time() - start_time) * 1000 # Calculate accurate cost using input/output split - cost = self.estimate_cost( - tokens_used, model, prompt_tokens=prompt_tokens, completion_tokens=completion_tokens + cost = self._calculate_response_cost( + model=model, + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + usage=data.get("usage") or {}, ) # ============================================================ @@ -1311,6 +1515,8 @@ async def _complete_impl( "finish_reason": choice["finish_reason"], "prompt_tokens": prompt_tokens, "completion_tokens": completion_tokens, + "cached_input_tokens": _cached_input_tokens(data.get("usage") or {}), + "cache_write_input_tokens": _cache_write_input_tokens(data.get("usage") or {}), # NEW: Add confidence analysis details for test validation "query": prompt, "confidence_method": confidence_method, @@ -1416,6 +1622,13 @@ async def _stream_impl( ... print(chunk, end='', flush=True) Python is a high-level programming language... """ + stable_cache_prefix = kwargs.pop("_cascadeflow_knowledge_cache_prefix", None) + explicit_cache = bool(stable_cache_prefix) and _uses_explicit_prompt_cache_breakpoints( + model + ) + if explicit_cache: + kwargs["prompt_cache_options"] = {"mode": "explicit"} + # Build messages messages = [] if system_prompt: @@ -1432,6 +1645,10 @@ async def _stream_impl( if self._should_use_responses_api(model): input_messages, instructions = self._convert_messages_to_responses_input(messages) + if explicit_cache: + input_messages, instructions = self._apply_responses_cache_breakpoint( + input_messages, instructions, stable_cache_prefix + ) payload: dict[str, Any] = { "model": model, "input": input_messages, @@ -1511,6 +1728,8 @@ async def _stream_impl( ) # Default: Chat Completions streaming. + if explicit_cache: + messages = self._apply_chat_cache_breakpoint(messages, stable_cache_prefix) payload: dict[str, Any] = { "model": model, "messages": messages, @@ -1606,43 +1825,7 @@ def estimate_cost( Returns: Estimated cost in USD """ - # OpenAI pricing per 1K tokens (as of January 2025) - # Source: https://openai.com/api/pricing/ - pricing = { - # GPT-5 series (current flagship - released August 2025) - # 50% cheaper input than GPT-4o, superior performance on coding, reasoning, math - "gpt-5": {"input": 0.00125, "output": 0.010}, - "gpt-5-mini": {"input": 0.00025, "output": 0.002}, - "gpt-5-nano": {"input": 0.00005, "output": 0.0004}, - "gpt-5-chat-latest": {"input": 0.00125, "output": 0.010}, - # GPT-4o series (previous flagship) - "gpt-4o": {"input": 0.0025, "output": 0.010}, - "gpt-4o-mini": {"input": 0.00015, "output": 0.0006}, - # O1 series (reasoning models) - "o1-preview": {"input": 0.015, "output": 0.060}, - "o1-mini": {"input": 0.003, "output": 0.012}, - "o1": {"input": 0.015, "output": 0.060}, - "o1-2024-12-17": {"input": 0.015, "output": 0.060}, - # O3 series (reasoning models) - "o3-mini": {"input": 0.001, "output": 0.005}, - # GPT-4 series (previous generation) - "gpt-4-turbo": {"input": 0.010, "output": 0.030}, - "gpt-4": {"input": 0.030, "output": 0.060}, - # GPT-3.5 series (deprecated - use gpt-4o-mini instead) - "gpt-3.5-turbo": {"input": 0.0005, "output": 0.0015}, - } - - # Find model pricing - model_pricing = None - model_lower = model.lower() - for prefix, rates in pricing.items(): - if model_lower.startswith(prefix): - model_pricing = rates - break - - # Default to GPT-4 pricing if unknown - if not model_pricing: - model_pricing = {"input": 0.030, "output": 0.060} + model_pricing = _openai_model_pricing(model) # Calculate accurate cost if we have the split if prompt_tokens is not None and completion_tokens is not None: diff --git a/cascadeflow/schema/usage.py b/cascadeflow/schema/usage.py index d6ec79ef..9e1cd495 100644 --- a/cascadeflow/schema/usage.py +++ b/cascadeflow/schema/usage.py @@ -11,6 +11,7 @@ class Usage: input_tokens: int = 0 output_tokens: int = 0 cached_input_tokens: int = 0 + cache_write_input_tokens: int = 0 @property def total_tokens(self) -> int: @@ -38,10 +39,15 @@ def from_payload(cls, usage: Any) -> "Usage": if cached_input_tokens is None: cached_input_tokens = usage.get("cache_read_input_tokens", 0) + cache_write_input_tokens = usage.get("cache_write_input_tokens") + if cache_write_input_tokens is None: + cache_write_input_tokens = usage.get("cache_creation_input_tokens", 0) + return cls( input_tokens=int(input_tokens or 0), output_tokens=int(output_tokens or 0), cached_input_tokens=int(cached_input_tokens or 0), + cache_write_input_tokens=int(cache_write_input_tokens or 0), ) def to_dict(self) -> dict[str, int]: @@ -49,5 +55,6 @@ def to_dict(self) -> dict[str, int]: "input_tokens": self.input_tokens, "output_tokens": self.output_tokens, "cached_input_tokens": self.cached_input_tokens, + "cache_write_input_tokens": self.cache_write_input_tokens, "total_tokens": self.total_tokens, } diff --git a/cascadeflow/streaming/base.py b/cascadeflow/streaming/base.py index 79500d64..10478741 100644 --- a/cascadeflow/streaming/base.py +++ b/cascadeflow/streaming/base.py @@ -23,6 +23,7 @@ from enum import Enum from typing import Any, Optional +from ..context import provider_cache_kwargs from ..utils.messages import messages_to_prompt logger = logging.getLogger(__name__) @@ -375,6 +376,7 @@ async def stream( StreamEvent objects with type, content, and data """ try: + prepared_knowledge = kwargs.pop("_cascadeflow_prepared_knowledge", None) query_text = messages_to_prompt(messages) if messages else query query = query_text logger.info(f"Starting streaming execution for query: {query_text[:50]}...") @@ -395,6 +397,14 @@ async def stream( "messages", } } + draft_provider_kwargs = { + **provider_kwargs, + **provider_cache_kwargs(self.cascade.drafter.provider, prepared_knowledge), + } + verifier_provider_kwargs = { + **provider_kwargs, + **provider_cache_kwargs(self.cascade.verifier.provider, prepared_knowledge), + } # Add tools and tool_choice if provided if tools is not None: @@ -405,7 +415,7 @@ async def stream( # ================================================================ # FIX #5: Add logprobs ONLY if no tools present (OpenAI limitation) # ================================================================ - logprobs_kwargs = provider_kwargs.copy() + logprobs_kwargs = draft_provider_kwargs.copy() provider_type = self.cascade.drafter.provider has_tools = tools is not None or "tools" in provider_kwargs @@ -440,7 +450,7 @@ async def stream( verifier_chunks = [] verifier_content = "" - verifier_logprobs_kwargs = provider_kwargs.copy() + verifier_logprobs_kwargs = verifier_provider_kwargs.copy() has_tools = tools is not None or "tools" in provider_kwargs if self.cascade.verifier.provider in ["openai"] and not has_tools: @@ -471,7 +481,7 @@ async def stream( prompt=query, max_tokens=max_tokens, temperature=temperature, - **provider_kwargs, + **verifier_provider_kwargs, ) verifier_content = response.content yield StreamEvent( @@ -605,7 +615,7 @@ async def stream( prompt=query, max_tokens=max_tokens, temperature=temperature, - **provider_kwargs, + **draft_provider_kwargs, ) draft_content = response.content draft_latency_ms = (time.time() - draft_start_time) * 1000 @@ -784,7 +794,7 @@ async def stream( verifier_chunks = [] verifier_content = "" - verifier_logprobs_kwargs = provider_kwargs.copy() + verifier_logprobs_kwargs = verifier_provider_kwargs.copy() has_tools = tools is not None or "tools" in provider_kwargs if self.cascade.verifier.provider in ["openai"] and not has_tools: @@ -826,7 +836,7 @@ async def stream( prompt=query, max_tokens=max_tokens, temperature=temperature, - **provider_kwargs, + **verifier_provider_kwargs, ) verifier_content = response.content verifier_latency_ms = (time.time() - verifier_start_time) * 1000 diff --git a/cascadeflow/streaming/tools.py b/cascadeflow/streaming/tools.py index a04b6a5b..aab3f639 100644 --- a/cascadeflow/streaming/tools.py +++ b/cascadeflow/streaming/tools.py @@ -43,6 +43,7 @@ from enum import Enum from typing import Any, Callable, Optional +from ..context import provider_cache_kwargs from ..utils.messages import messages_to_prompt, normalize_messages from .utils import ( JSONParseState, @@ -432,6 +433,7 @@ async def stream( raise ValueError("tools parameter is required for tool streaming") try: + prepared_knowledge = kwargs.pop("_cascadeflow_prepared_knowledge", None) normalized_messages = normalize_messages(messages) if messages else None query_text = messages_to_prompt(normalized_messages) if normalized_messages else query @@ -492,6 +494,10 @@ async def stream( draft_provider = self.cascade.providers[self.cascade.drafter.provider] draft_model = self.cascade.drafter + draft_provider_kwargs = { + **provider_kwargs, + **provider_cache_kwargs(draft_model.provider, prepared_knowledge), + } draft_chunks = [] draft_content = "" @@ -523,7 +529,7 @@ def _provider_supports_tools(p: Any) -> bool: max_tokens=max_tokens, temperature=temperature, tool_choice=tool_choice, # ← Explicit - **provider_kwargs, # ← Does NOT contain tools/tool_choice + **draft_provider_kwargs, # ← Does NOT contain tools/tool_choice ): # Process chunk for tool calls async for event in self._process_tool_chunk( @@ -562,7 +568,7 @@ def _provider_supports_tools(p: Any) -> bool: max_tokens=max_tokens, temperature=temperature, tool_choice=tool_choice, # ← Explicit - **provider_kwargs, # ← Does NOT contain tools/tool_choice + **draft_provider_kwargs, # ← Does NOT contain tools/tool_choice ) # Extract tool calls from response @@ -863,7 +869,7 @@ def _provider_supports_tools(p: Any) -> bool: max_tokens=max_tokens, temperature=temperature, tool_choice=tool_choice, - **provider_kwargs, + **draft_provider_kwargs, ) next_tool_calls = ( @@ -1006,6 +1012,10 @@ def _provider_supports_tools(p: Any) -> bool: verifier_start_time = time.time() verifier_provider = self.cascade.providers[self.cascade.verifier.provider] verifier_model = self.cascade.verifier + verifier_provider_kwargs = { + **provider_kwargs, + **provider_cache_kwargs(verifier_model.provider, prepared_knowledge), + } verifier_tool_calls = [] verifier_content = "" @@ -1024,7 +1034,7 @@ def _provider_supports_tools(p: Any) -> bool: max_tokens=max_tokens, temperature=temperature, tool_choice=tool_choice, # ← Explicit - **provider_kwargs, # ← Clean kwargs + **verifier_provider_kwargs, # ← Clean kwargs ): # Process chunk async for event in self._process_tool_chunk(chunk, tools, None, ""): @@ -1058,7 +1068,7 @@ def _provider_supports_tools(p: Any) -> bool: max_tokens=max_tokens, temperature=temperature, tool_choice=tool_choice, # ← Explicit - **provider_kwargs, # ← Clean kwargs + **verifier_provider_kwargs, # ← Clean kwargs ) if hasattr(response, "tool_calls") and response.tool_calls: @@ -1156,7 +1166,7 @@ def _provider_supports_tools(p: Any) -> bool: max_tokens=max_tokens, temperature=temperature, tool_choice=tool_choice, - **provider_kwargs, + **verifier_provider_kwargs, ) next_tool_calls = ( diff --git a/docs/README.md b/docs/README.md index 6553688e..f42ed8e7 100644 --- a/docs/README.md +++ b/docs/README.md @@ -25,6 +25,7 @@ Agent runtime intelligence layer — optimize cost, latency, quality, budget, co - [Agentic Patterns (TypeScript)](guides/agentic-typescript.md) - Tool loops, multi-agent orchestration, and message best practices - [Harness Telemetry & Privacy](guides/harness_telemetry_privacy.md) - Decision traces, callbacks, and privacy-safe observability - [Cost Tracking](guides/cost_tracking.md) - Track and analyze API costs across queries +- [Versioned Knowledge Context](guides/knowledge_context.md) - Safe, cache-aware knowledge handoff across model switches - [Proxy Routing](guides/proxy.md) - Route requests through provider-aware proxy plans ## 🏭 Production & Advanced @@ -46,6 +47,7 @@ Agent runtime intelligence layer — optimize cost, latency, quality, budget, co - [Google ADK Integration](guides/google_adk_integration.md) - Plugin-based harness integration for ADK runners (opt-in) - [n8n Integration](guides/n8n_integration.md) - Use cascadeflow in n8n workflows - [Paygentic Integration](guides/paygentic_integration.md) - Usage metering and billing lifecycle helpers (opt-in) +- [MCP Integration](guides/mcp_integration.md) - Use cascadeflow routing from ChatGPT and Claude clients ## 📚 Examples diff --git a/docs/guides/knowledge_context.md b/docs/guides/knowledge_context.md new file mode 100644 index 00000000..66a03850 --- /dev/null +++ b/docs/guides/knowledge_context.md @@ -0,0 +1,96 @@ +# Versioned knowledge across model switches + +Cascadeflow can attach one immutable knowledge snapshot to a request and pass the +same snapshot to every model selected during direct, draft, verifier, tool, and +streaming execution. Knowledge selection is request-scoped: the agent never keeps +an implicit "active" knowledge set that a concurrent or later request could inherit. + +## Python + +```python +from cascadeflow import CascadeAgent, KnowledgeSnapshot + +knowledge = KnowledgeSnapshot( + key="support-manual", + version="2026-08-06", + content=retrieved_text, + cache_ttl="1h", # "5m" is the default +) + +result = await agent.run("How do I reset the device?", knowledge=knowledge) +print(result.metadata["knowledge"]) +``` + +Pass a different snapshot on the next request to switch knowledge. Reusing the +same `key` and explicit `version` with different content is rejected while that +version is resident locally. Provider cache keys also include the content digest, +so old provider-side content cannot be selected after an LRU eviction or restart. + +## TypeScript + +```ts +const result = await agent.run('How do I reset the device?', { + knowledge: { + key: 'support-manual', + version: '2026-08-06', + content: retrievedText, + cacheTtl: '1h', + }, +}); +``` + +The same option works with `runStream`. + +## Cost model + +The local LRU stores only the stable rendered prefix and metadata. It does **not** +pretend to remove provider input tokens. Correctness always comes from sending the +selected snapshot to each model that handles the request. + +To reduce billed work without risking stale state: + +1. Retrieve only the passages needed for the current request before constructing + the snapshot. Do not attach an entire knowledge base by default. +2. Keep knowledge and stable instructions first, and the current conversation turn + last. Provider prompt caches require a stable prefix. +3. Reuse the same key/version/content while the knowledge is unchanged. Change the + version when it changes. +4. Send conversation history only when the query depends on it. Prefer a concise, + application-owned handoff over blindly replaying a full transcript. +5. Measure `cached_input_tokens` and `cache_write_input_tokens` in provider metadata + before claiming savings. CascadeFlow uses those reported categories in OpenAI + and Anthropic result-cost calculations instead of assuming every request hit. + +OpenAI receives a content-bound `prompt_cache_key`. For GPT-5.6 and later model +families, cascadeflow also places an explicit breakpoint immediately after the +stable knowledge prefix and disables the implicit latest-message breakpoint. This +prevents each changing query from creating a separate billable cache write. +Anthropic receives top-level `cache_control` with a five-minute or one-hour TTL. +Other providers—including local Ollama and vLLM endpoints—receive the same +knowledge prefix without unsupported cache arguments. Provider caching is therefore +an optimization, never a dependency. + +Provider caches have minimum prefix sizes and charge for the initial write on some +models. A short or one-off snapshot may not save money. Disable provider caching +with `enable_provider_cache=False` (Python) or `enableProviderCache: false` +(TypeScript) when a snapshot is not expected to be reused. The local immutable +snapshot and switch-safety behavior remain enabled. + +Cache-aware dollar estimates use standard API rates. Provider service tiers, +regional processing, cloud resellers, negotiated pricing, and future price changes +can differ, so provider invoices remain authoritative. + +See the official [OpenAI prompt caching guide](https://developers.openai.com/api/docs/guides/prompt-caching) +and [Anthropic prompt caching guide](https://platform.claude.com/docs/en/build-with-claude/prompt-caching) +for provider-specific cache thresholds, billing, and retention behavior. + +## Concurrency and security + +- The snapshot is immutable and selected per call, so concurrent tenants cannot + race through shared "current knowledge" state. +- The provider cache key hashes the logical identity and content digest; it does not + expose tenant or document names. +- Knowledge text is still sent to the selected provider models. Apply the same data + classification and provider policy used for normal prompts. +- A `key` and `version` are identifiers, not authorization. Resolve tenant access + before constructing a snapshot. diff --git a/docs/guides/mcp_integration.md b/docs/guides/mcp_integration.md new file mode 100644 index 00000000..f60dd5d2 --- /dev/null +++ b/docs/guides/mcp_integration.md @@ -0,0 +1,87 @@ +# ChatGPT and Claude integration through MCP + +Yes: one cascadeflow MCP server can expose cost-aware routing as a tool to ChatGPT, +Claude, Claude Desktop, and other MCP clients. The initial integration is tool-first +and deliberately has no required UI. + +```text +ChatGPT / Claude + -> small MCP tool call (query + knowledge ID + optional concise handoff) + -> cascadeflow MCP server + -> server-side knowledge resolver + -> draft / verifier / direct provider calls + <- answer + compact routing and cost trace +``` + +This brings cascadeflow logic into a conversation, but it does not replace the host +application's own model. The host model still decides to call the tool, and +cascadeflow then makes its provider calls. That extra host turn and latency must be +included in the end-to-end economics. + +## Server factory + +Install the optional SDK on Python 3.10+: + +```bash +pip install "cascadeflow[mcp]" +``` + +Create a configured agent and keep knowledge lookup behind the MCP server: + +```python +from cascadeflow import CascadeAgent, KnowledgeSnapshot +from cascadeflow.integrations.mcp import create_mcp_server + +async def resolve_knowledge(key: str, version: str | None): + content, resolved_version = await knowledge_store.load(key, version) + return KnowledgeSnapshot( + key=key, + version=resolved_version, + content=content, + cache_ttl="1h", + ) + +mcp = create_mcp_server(agent, knowledge_resolver=resolve_knowledge) + +# Local desktop transport: +mcp.run(transport="stdio") + +# Remote ChatGPT/Claude transport: +# mcp.run(transport="streamable-http") +``` + +The `cascadeflow_run` tool accepts: + +- `query`: the current request; +- `knowledge_key` and optional `knowledge_version`: server-side selection only; +- `conversation_context`: an optional concise factual handoff when prior turns are + genuinely required. It is bounded by default to prevent accidental transcript + replay. + +The tool returns the answer plus a compact model, routing, cost, latency, and +knowledge-version trace. It never returns the private knowledge content. + +## Host choices + +- **ChatGPT:** deploy a public HTTPS Streamable HTTP endpoint (normally `/mcp`) and + add it as a custom MCP connection in developer mode. OpenAI's current plugin + documentation describes [building the MCP server](https://developers.openai.com/plugins/build/mcp-server) + and [connecting it to ChatGPT](https://developers.openai.com/plugins/deploy/connect-chatgpt). +- **Claude and Claude Desktop:** use the same public Streamable HTTP endpoint as a + remote custom connector. Claude Desktop can also package a local stdio server as + a desktop extension. See Anthropic's [remote connector guide](https://support.claude.com/en/articles/11175166-get-started-with-custom-connectors-using-remote-mcp) + and [desktop extension guide](https://support.claude.com/en/articles/10949351-getting-started-with-local-mcp-servers-on-claude-desktop). + +Production remote deployments should use authentication, tenant-scoped knowledge +authorization, rate limits, audit logs, TLS, and a public endpoint reachable by the +host. Do not pass provider API keys as MCP tool arguments. + +## MCP Apps / interactive UI + +An interactive routing panel is a useful second layer, not part of the routing +critical path. The standard [MCP Apps extension](https://modelcontextprotocol.io/extensions/apps/overview) +can render a sandboxed `ui://` resource for model choice, knowledge-version display, +and savings/latency traces. Keep `cascadeflow_run` fully usable without that UI so +tool behavior remains portable across hosts. If a host needs product-specific UI +metadata, add a thin host adapter around the same tool and server-side logic rather +than forking cascadeflow's routing implementation. diff --git a/packages/core/src/__tests__/knowledge-cache.test.ts b/packages/core/src/__tests__/knowledge-cache.test.ts new file mode 100644 index 00000000..28fc6722 --- /dev/null +++ b/packages/core/src/__tests__/knowledge-cache.test.ts @@ -0,0 +1,157 @@ +import { beforeAll, describe, expect, it, vi } from 'vitest'; + +import { CascadeAgent } from '../agent'; +import type { ModelConfig } from '../config'; +import { + KnowledgeCache, + providerKnowledgeCacheOptions, +} from '../knowledge-cache'; +import { providerRegistry, type Provider, type ProviderRequest } from '../providers/base'; +import { OpenAIProvider } from '../providers/openai'; +import type { ProviderResponse } from '../types'; + +const capturedRequests: ProviderRequest[] = []; + +class KnowledgeCaptureProvider implements Provider { + readonly name = 'knowledge-capture'; + + constructor(_config: ModelConfig) {} + + async generate(request: ProviderRequest): Promise { + capturedRequests.push(request); + return { + content: 'ok', + model: request.model, + usage: { prompt_tokens: 1, completion_tokens: 1, total_tokens: 2 }, + }; + } + + calculateCost(): number { + return 0; + } + + isAvailable(): boolean { + return true; + } +} + +beforeAll(() => { + providerRegistry.register('knowledge-capture' as any, KnowledgeCaptureProvider as any); +}); + +describe('KnowledgeCache', () => { + it('reuses an immutable version and records a local cache hit', () => { + const cache = new KnowledgeCache(2); + const snapshot = { key: 'catalog', version: '2026-08-06', content: 'Current catalog' }; + + const first = cache.prepare(snapshot); + const second = cache.prepare(snapshot); + + expect(first.localCacheHit).toBe(false); + expect(second.localCacheHit).toBe(true); + expect(second.prepared).toBe(first.prepared); + expect(cache.getStats()).toEqual({ size: 1, hits: 1, misses: 1, evictions: 0 }); + }); + + it('rejects reusing a version with different content', () => { + const cache = new KnowledgeCache(); + cache.prepare({ key: 'policy', version: 'v1', content: 'Policy A' }); + + expect(() => cache.prepare({ key: 'policy', version: 'v1', content: 'Policy B' })).toThrow( + /version collision/ + ); + }); + + it('keeps provider keys content-bound after local eviction', () => { + const cache = new KnowledgeCache(1); + const old = cache.prepare({ key: 'docs', version: 'v1', content: 'old' }).prepared; + cache.prepare({ key: 'other', version: 'v1', content: 'other' }); + const fresh = cache.prepare({ key: 'docs', version: 'v1', content: 'new' }).prepared; + + expect(fresh.promptCacheKey).not.toBe(old.promptCacheKey); + }); + + it('emits native cache hints only for providers that support them', () => { + const prepared = new KnowledgeCache().prepare({ + key: 'manual', + version: 'v3', + content: 'Manual text', + cacheTtl: '1h', + }).prepared; + + expect(providerKnowledgeCacheOptions('openai', prepared)).toEqual({ + prompt_cache_key: prepared.promptCacheKey, + _cascadeflow_knowledge_cache_prefix: prepared.systemPrefix, + }); + expect(providerKnowledgeCacheOptions('anthropic', prepared)).toEqual({ + cache_control: { type: 'ephemeral', ttl: '1h' }, + }); + expect(providerKnowledgeCacheOptions('ollama', prepared)).toEqual({}); + }); + + it('switches snapshots without leaking previous knowledge into a request', async () => { + capturedRequests.length = 0; + const agent = new CascadeAgent({ + models: [{ name: 'capture', provider: 'knowledge-capture' as any, cost: 0 }], + }); + + await agent.run('question', { + systemPrompt: 'Answer briefly.', + knowledge: { key: 'tenant', version: 'alpha', content: 'Alpha facts' }, + }); + await agent.run('question', { + systemPrompt: 'Answer briefly.', + knowledge: { key: 'tenant', version: 'beta', content: 'Beta facts' }, + }); + + expect(capturedRequests).toHaveLength(2); + expect(capturedRequests[0].systemPrompt).toContain('Alpha facts'); + expect(capturedRequests[0].systemPrompt).not.toContain('Beta facts'); + expect(capturedRequests[1].systemPrompt).toContain('Beta facts'); + expect(capturedRequests[1].systemPrompt).not.toContain('Alpha facts'); + expect(capturedRequests[1].systemPrompt).toContain('Answer briefly.'); + }); + + it('places a GPT-5.6 breakpoint after stable knowledge without leaking internal fields', async () => { + const snapshot = new KnowledgeCache().prepare({ + key: 'manual', + version: 'v1', + content: 'Stable manual content', + }).prepared; + const create = vi.fn().mockResolvedValue({ + choices: [{ message: { content: 'ok' }, finish_reason: 'stop' }], + model: 'gpt-5.6-terra', + usage: { + prompt_tokens: 1100, + completion_tokens: 1, + total_tokens: 1101, + prompt_tokens_details: { cached_tokens: 1024 }, + cache_write_tokens: 0, + }, + }); + const provider = new OpenAIProvider({ + name: 'gpt-5.6-terra', + provider: 'openai', + cost: 0, + apiKey: 'test', + }); + (provider as any).useSDK = true; + (provider as any).client = { chat: { completions: { create } } }; + + const extra = providerKnowledgeCacheOptions('openai', snapshot); + const result = await provider.generate({ + model: 'gpt-5.6-terra', + messages: [{ role: 'user', content: 'Changing question' }], + systemPrompt: snapshot.systemPrefix, + extra, + }); + + const payload = create.mock.calls[0][0]; + expect(payload.prompt_cache_options).toEqual({ mode: 'explicit' }); + expect(payload.messages[0].content[0].prompt_cache_breakpoint).toEqual({ mode: 'explicit' }); + expect(payload.messages[0].content[0].text).toBe(snapshot.systemPrefix); + expect(JSON.stringify(payload)).not.toContain('_cascadeflow_knowledge_cache_prefix'); + expect(result.usage?.cached_input_tokens).toBe(1024); + expect(result.usage?.cache_write_input_tokens).toBe(0); + }); +}); diff --git a/packages/core/src/__tests__/reasoning-models.test.ts b/packages/core/src/__tests__/reasoning-models.test.ts index 2aefa6f0..a4c2bcbe 100644 --- a/packages/core/src/__tests__/reasoning-models.test.ts +++ b/packages/core/src/__tests__/reasoning-models.test.ts @@ -166,6 +166,11 @@ describe('Reasoning Model Support', () => { describe('Model-specific Pricing', () => { const testCases = [ + // GPT-5.6 series + { model: 'gpt-5.6', input: 0.005, output: 0.030 }, + { model: 'gpt-5.6-terra', input: 0.002, output: 0.012 }, + { model: 'gpt-5.6-luna', input: 0.0002, output: 0.0012 }, + // GPT-5 series { model: 'gpt-5', input: 0.00125, output: 0.010 }, { model: 'gpt-5-mini', input: 0.00025, output: 0.002 }, @@ -201,6 +206,31 @@ describe('Reasoning Model Support', () => { expect(cost).toBeCloseTo(expectedCost, 6); }); }); + + it('prices GPT-5.6 cache reads and writes from reported categories', () => { + const provider = new OpenAIProvider({ ...mockConfig, name: 'gpt-5.6-terra' }); + const hitCost = provider.calculateCostFromUsage( + { + prompt_tokens: 1100, + completion_tokens: 1, + total_tokens: 1101, + cached_input_tokens: 1024, + }, + 'gpt-5.6-terra' + ); + const writeCost = provider.calculateCostFromUsage( + { + prompt_tokens: 1100, + completion_tokens: 1, + total_tokens: 1101, + cache_write_input_tokens: 1024, + }, + 'openai/gpt-5.6-terra' + ); + + expect(hitCost).toBeCloseTo(0.0003688, 8); + expect(writeCost).toBeCloseTo(0.002724, 8); + }); }); describe('Edge Cases', () => { @@ -439,7 +469,14 @@ describe('Anthropic Claude 4.5 Extended Thinking', () => { describe('Anthropic Pricing Matrix', () => { const testCases = [ + // Claude 5 Series + { model: 'claude-fable-5', expectedBlended: 30.0 }, + { model: 'claude-opus-5', expectedBlended: 15.0 }, + { model: 'claude-sonnet-5', expectedBlended: 6.0 }, + // Claude 4 Series + { model: 'claude-opus-4-8', expectedBlended: 15.0 }, + { model: 'claude-sonnet-4-6', expectedBlended: 9.0 }, { model: 'claude-opus-4.1', expectedBlended: 45.0 }, { model: 'claude-opus-4', expectedBlended: 45.0 }, { model: 'claude-sonnet-4.5', expectedBlended: 9.0 }, @@ -447,7 +484,7 @@ describe('Anthropic Claude 4.5 Extended Thinking', () => { // Claude 3.5 Series { model: 'claude-3-5-sonnet', expectedBlended: 9.0 }, - { model: 'claude-3-5-haiku', expectedBlended: 3.0 }, + { model: 'claude-3-5-haiku', expectedBlended: 2.4 }, // Claude 3 Series { model: 'claude-3-opus', expectedBlended: 45.0 }, @@ -465,6 +502,23 @@ describe('Anthropic Claude 4.5 Extended Thinking', () => { expect(cost).toBeCloseTo(expectedCost, 6); }); }); + + it('prices cache reads and one-hour writes separately', () => { + const provider = new AnthropicProvider({ ...mockConfig, name: 'claude-sonnet-4-5' }); + const cost = provider.calculateCostFromUsage( + { + prompt_tokens: 100, + completion_tokens: 10, + total_tokens: 110, + cached_input_tokens: 1000, + cache_write_input_tokens: 1000, + cache_write_1h_input_tokens: 1000, + }, + 'claude-sonnet-4-5' + ); + + expect(cost).toBeCloseTo(0.00675, 8); + }); }); describe('Edge Cases', () => { diff --git a/packages/core/src/__tests__/types-usage.test.ts b/packages/core/src/__tests__/types-usage.test.ts index 95eb482a..02694e60 100644 --- a/packages/core/src/__tests__/types-usage.test.ts +++ b/packages/core/src/__tests__/types-usage.test.ts @@ -7,14 +7,21 @@ describe('canonical usage mapping', () => { const usage = toCanonicalUsage({ prompt_tokens: 10, completion_tokens: 15, total_tokens: 25 }); expect(usage.input_tokens).toBe(10); expect(usage.output_tokens).toBe(15); + expect(usage.cache_write_input_tokens).toBe(0); expect(usage.total_tokens).toBe(25); }); it('preserves canonical fields', () => { - const usage = toCanonicalUsage({ input_tokens: 8, output_tokens: 12, cached_input_tokens: 4 } as any); + const usage = toCanonicalUsage({ + input_tokens: 8, + output_tokens: 12, + cached_input_tokens: 4, + cache_write_input_tokens: 3, + } as any); expect(usage.input_tokens).toBe(8); expect(usage.output_tokens).toBe(12); expect(usage.cached_input_tokens).toBe(4); + expect(usage.cache_write_input_tokens).toBe(3); expect(usage.total_tokens).toBe(20); }); }); diff --git a/packages/core/src/agent.ts b/packages/core/src/agent.ts index a7d224fb..b4bd38a8 100644 --- a/packages/core/src/agent.ts +++ b/packages/core/src/agent.ts @@ -2,7 +2,7 @@ * cascadeflow Agent - MVP Implementation */ -import { providerRegistry, getAvailableProviders } from './providers/base'; +import { providerRegistry, getAvailableProviders, type Provider } from './providers/base'; import { OpenAIProvider } from './providers/openai'; import { AnthropicProvider } from './providers/anthropic'; import { GroqProvider } from './providers/groq'; @@ -14,7 +14,7 @@ import { OpenRouterProvider } from './providers/openrouter'; import { VercelAISDKProvider, VERCEL_AI_PROVIDER_NAMES } from './providers/vercel-ai'; import type { AgentConfig, ModelConfig } from './config'; import type { CascadeResult } from './result'; -import type { Message, Tool, UserProfile, TierLevel } from './types'; +import type { Message, Tool, UserProfile, TierLevel, ProviderResponse } from './types'; import { ToolCall as ParsedToolCall, ToolExecutor } from './tools'; import { type StreamEvent, @@ -41,6 +41,12 @@ import { CallbackEvent } from './telemetry/callbacks'; import type { DomainConfig, DomainConfigMap } from './config/domain-config'; import { RuleEngine } from './rules'; import type { RuleContext, RuleDecision } from './rules'; +import { + KnowledgeCache, + providerKnowledgeCacheOptions, + type KnowledgeSnapshot, + type PreparedKnowledge, +} from './knowledge-cache'; // Register providers providerRegistry.register('openai', OpenAIProvider); @@ -98,6 +104,9 @@ export interface RunOptions { /** System prompt */ systemPrompt?: string; + /** Immutable request-scoped knowledge used by every routed model. */ + knowledge?: string | KnowledgeSnapshot; + /** Tools/functions available */ tools?: Tool[]; @@ -169,6 +178,29 @@ export class CascadeAgent { private enableDomainDetection: boolean; private ruleEngine: RuleEngine; private toolExecutor?: ToolExecutor; + private knowledgeCache = new KnowledgeCache(); + + private providerExtra( + extra: Record | undefined, + provider: string, + prepared?: PreparedKnowledge + ): Record | undefined { + const cacheOptions = providerKnowledgeCacheOptions(provider, prepared); + if (!extra && Object.keys(cacheOptions).length === 0) return undefined; + return { ...(extra ?? {}), ...cacheOptions }; + } + + private responseCost(provider: Provider, response: ProviderResponse): number { + if (!response.usage) return 0; + if (provider.calculateCostFromUsage) { + return provider.calculateCostFromUsage(response.usage, response.model); + } + return provider.calculateCost( + response.usage.prompt_tokens, + response.usage.completion_tokens, + response.model + ); + } /** * Create a new cascadeflow agent @@ -735,7 +767,13 @@ export class CascadeAgent { const rawMessages: Message[] = typeof input === 'string' ? [{ role: 'user', content: input }] : input; - const normalized = normalizeSystemPromptFromMessages(rawMessages, options.systemPrompt); + const preparedKnowledge = options.knowledge + ? this.knowledgeCache.prepare(options.knowledge).prepared + : undefined; + const stableSystemPrompt = preparedKnowledge + ? [preparedKnowledge.systemPrefix, options.systemPrompt].filter(Boolean).join('\n\n') + : options.systemPrompt; + const normalized = normalizeSystemPromptFromMessages(rawMessages, stableSystemPrompt); const messages = normalized.messages; const executor = options.toolExecutor ?? this.toolExecutor; if (executor && typeof (executor as any).executeParallel !== 'function') { @@ -914,7 +952,7 @@ export class CascadeAgent { temperature: options.temperature, systemPrompt: normalized.systemPrompt, tools: options.tools, - extra: options.extra, + extra: this.providerExtra(options.extra, bestModelConfig.provider, preparedKnowledge), }); modelUsed = response.model; @@ -922,11 +960,7 @@ export class CascadeAgent { finalToolCalls = response.tool_calls; if (response.usage) { - totalLoopCost += provider.calculateCost( - response.usage.prompt_tokens, - response.usage.completion_tokens, - response.model - ); + totalLoopCost += this.responseCost(provider, response); } const assistantMsg: Message = { role: 'assistant', content: response.content || '' }; @@ -991,7 +1025,7 @@ export class CascadeAgent { temperature: options.temperature, systemPrompt: normalized.systemPrompt, tools: options.tools, - extra: options.extra, + extra: this.providerExtra(options.extra, bestModelConfig.provider, preparedKnowledge), }); modelUsed = response.model; @@ -1000,11 +1034,7 @@ export class CascadeAgent { // Calculate cost if (response.usage) { - totalCost = provider.calculateCost( - response.usage.prompt_tokens, - response.usage.completion_tokens, - response.model - ); + totalCost = this.responseCost(provider, response); } const latencyMs = Date.now() - startTime; @@ -1051,7 +1081,7 @@ export class CascadeAgent { temperature: options.temperature, systemPrompt: normalized.systemPrompt, tools: options.tools, - extra: options.extra, + extra: this.providerExtra(options.extra, draftModelConfig.provider, preparedKnowledge), }); draftLatency = Date.now() - draftStart; @@ -1062,11 +1092,7 @@ export class CascadeAgent { // Calculate draft cost if (draftResponse.usage) { - draftCost = draftProvider.calculateCost( - draftResponse.usage.prompt_tokens, - draftResponse.usage.completion_tokens, - draftResponse.model - ); + draftCost = this.responseCost(draftProvider, draftResponse); } // Quality validation using logprobs and heuristics @@ -1141,7 +1167,7 @@ export class CascadeAgent { temperature: options.temperature, systemPrompt: normalized.systemPrompt, tools: options.tools, - extra: options.extra, + extra: this.providerExtra(options.extra, verifierModelConfig.provider, preparedKnowledge), }); verifierLatency = Date.now() - verifierStart; @@ -1152,11 +1178,7 @@ export class CascadeAgent { // Calculate verifier cost if (verifierResponse.usage) { - verifierCost = verifierProvider.calculateCost( - verifierResponse.usage.prompt_tokens, - verifierResponse.usage.completion_tokens, - verifierResponse.model - ); + verifierCost = this.responseCost(verifierProvider, verifierResponse); } } else { draftAccepted = true; @@ -1327,7 +1349,13 @@ export class CascadeAgent { const rawMessages: Message[] = typeof input === 'string' ? [{ role: 'user', content: input }] : input; - const normalized = normalizeSystemPromptFromMessages(rawMessages, options.systemPrompt); + const preparedKnowledge = options.knowledge + ? this.knowledgeCache.prepare(options.knowledge).prepared + : undefined; + const stableSystemPrompt = preparedKnowledge + ? [preparedKnowledge.systemPrefix, options.systemPrompt].filter(Boolean).join('\n\n') + : options.systemPrompt; + const normalized = normalizeSystemPromptFromMessages(rawMessages, stableSystemPrompt); const messages = normalized.messages; // Extract query text for complexity detection (exclude system messages) @@ -1481,19 +1509,20 @@ export class CascadeAgent { if (!directProvider.stream) { // Fallback to non-streaming path - const result = await this.run(input, { - maxTokens: maxTokens, - temperature: options.temperature, - systemPrompt: normalized.systemPrompt, - tools: options.tools, - extra: options.extra, - forceDirect: true, - userTier: options.userTier, - workflow: options.workflow, - kpiFlags: options.kpiFlags, - tenantId: options.tenantId, - channel: options.channel, - }); + const result = await this.run(input, { + maxTokens: maxTokens, + temperature: options.temperature, + systemPrompt: options.systemPrompt, + knowledge: options.knowledge, + tools: options.tools, + extra: options.extra, + forceDirect: true, + userTier: options.userTier, + workflow: options.workflow, + kpiFlags: options.kpiFlags, + tenantId: options.tenantId, + channel: options.channel, + }); yield createStreamEvent(StreamEventType.CHUNK, result.content, { model: result.modelUsed, phase: 'direct', @@ -1514,7 +1543,7 @@ export class CascadeAgent { temperature: options.temperature, systemPrompt: normalized.systemPrompt, tools: options.tools, - extra: options.extra, + extra: this.providerExtra(options.extra, bestModelConfig.provider, preparedKnowledge), })) { directContent += chunk.content; @@ -1595,7 +1624,7 @@ export class CascadeAgent { temperature: options.temperature, systemPrompt: normalized.systemPrompt, tools: options.tools, - extra: options.extra, + extra: this.providerExtra(options.extra, draftModelConfig.provider, preparedKnowledge), })) { draftContent += chunk.content; @@ -1737,7 +1766,7 @@ export class CascadeAgent { temperature: options.temperature, systemPrompt: normalized.systemPrompt, tools: options.tools, - extra: options.extra, + extra: this.providerExtra(options.extra, verifierModelConfig.provider, preparedKnowledge), })) { verifierContent += chunk.content; @@ -1770,7 +1799,7 @@ export class CascadeAgent { temperature: options.temperature, systemPrompt: normalized.systemPrompt, tools: options.tools, - extra: options.extra, + extra: this.providerExtra(options.extra, verifierModelConfig.provider, preparedKnowledge), }); verifierContent = verifierResponse.content; verifierToolCalls = verifierResponse.tool_calls; diff --git a/packages/core/src/index.ts b/packages/core/src/index.ts index c919f67e..aca8483b 100644 --- a/packages/core/src/index.ts +++ b/packages/core/src/index.ts @@ -71,6 +71,14 @@ export { export type { CascadeResult } from './result'; export { resultToObject } from './result'; +// Provider-neutral knowledge handoff +export { KnowledgeCache, providerKnowledgeCacheOptions } from './knowledge-cache'; +export type { + KnowledgeSnapshot, + PreparedKnowledge, + KnowledgeCacheTtl, +} from './knowledge-cache'; + // Batch Processing (v0.2.1+) export { BatchStrategy, diff --git a/packages/core/src/knowledge-cache.ts b/packages/core/src/knowledge-cache.ts new file mode 100644 index 00000000..b0929f50 --- /dev/null +++ b/packages/core/src/knowledge-cache.ts @@ -0,0 +1,145 @@ +/** Provider-neutral, versioned knowledge handoff for cross-model routing. */ + +import type { Provider } from './types'; + +export type KnowledgeCacheTtl = '5m' | '1h'; + +export interface KnowledgeSnapshot { + content: string; + key?: string; + version?: string; + enableProviderCache?: boolean; + cacheTtl?: KnowledgeCacheTtl; +} + +export interface PreparedKnowledge { + identity: string; + contentDigest: string; + systemPrefix: string; + enableProviderCache: boolean; + cacheTtl: KnowledgeCacheTtl; + promptCacheKey: string; +} + +function fingerprint(value: string): string { + // Two seeded FNV-1a passes keep this browser-safe and deterministic. This is + // an identity fingerprint, not a security primitive. + const pass = (seed: number): string => { + let hash = seed >>> 0; + for (let i = 0; i < value.length; i++) { + hash ^= value.charCodeAt(i); + hash = Math.imul(hash, 0x01000193) >>> 0; + } + return hash.toString(16).padStart(8, '0'); + }; + return `${pass(0x811c9dc5)}${pass(0x9e3779b9)}`; +} + +function safeLabel(value: string): string { + return value.trim().replace(/[^a-zA-Z0-9._:/-]+/g, '-').replace(/^-+|-+$/g, '') || 'knowledge'; +} + +export class KnowledgeCache { + private entries = new Map(); + private versionDigests = new Map(); + private hits = 0; + private misses = 0; + private evictions = 0; + + constructor(private readonly maxEntries = 128) { + if (maxEntries <= 0) throw new Error('maxEntries must be positive'); + } + + prepare(value: string | KnowledgeSnapshot): { prepared: PreparedKnowledge; localCacheHit: boolean } { + const snapshot: KnowledgeSnapshot = typeof value === 'string' ? { content: value } : value; + if (!snapshot.content?.trim()) throw new Error('KnowledgeSnapshot.content must be non-empty'); + if (snapshot.cacheTtl && snapshot.cacheTtl !== '5m' && snapshot.cacheTtl !== '1h') { + throw new Error("KnowledgeSnapshot.cacheTtl must be '5m' or '1h'"); + } + + const key = safeLabel(snapshot.key ?? 'default'); + const contentDigest = fingerprint(snapshot.content); + const version = safeLabel(snapshot.version ?? contentDigest); + const identity = `${key}:${version}`; + const enableProviderCache = snapshot.enableProviderCache ?? true; + const cacheTtl = snapshot.cacheTtl ?? '5m'; + const entryKey = JSON.stringify([identity, contentDigest, enableProviderCache, cacheTtl]); + const priorDigest = this.versionDigests.get(identity); + if (priorDigest && priorDigest !== contentDigest) { + throw new Error( + 'Knowledge version collision: reuse of a key/version with different content would risk stale context' + ); + } + + const existing = this.entries.get(entryKey); + if (existing) { + this.entries.delete(entryKey); + this.entries.set(entryKey, existing); + this.hits++; + return { prepared: existing, localCacheHit: true }; + } + + const prepared: PreparedKnowledge = { + identity, + contentDigest, + systemPrefix: + `\n` + + `${snapshot.content.trim()}\n`, + enableProviderCache, + cacheTtl, + promptCacheKey: `cascadeflow:knowledge:${fingerprint(`${identity}:${contentDigest}`)}`, + }; + this.entries.set(entryKey, prepared); + this.versionDigests.set(identity, contentDigest); + this.misses++; + + if (this.entries.size > this.maxEntries) { + const oldest = this.entries.keys().next().value as string | undefined; + if (oldest) { + const evictedIdentity = this.entries.get(oldest)?.identity; + this.entries.delete(oldest); + if ( + evictedIdentity && + !Array.from(this.entries.values()).some(entry => entry.identity === evictedIdentity) + ) { + this.versionDigests.delete(evictedIdentity); + } + this.evictions++; + } + } + return { prepared, localCacheHit: false }; + } + + clear(): void { + this.entries.clear(); + this.versionDigests.clear(); + } + + getStats(): { size: number; hits: number; misses: number; evictions: number } { + return { size: this.entries.size, hits: this.hits, misses: this.misses, evictions: this.evictions }; + } +} + +export function providerKnowledgeCacheOptions( + provider: Provider | string, + prepared?: PreparedKnowledge +): Record { + if (!prepared?.enableProviderCache) return {}; + if (provider === 'openai') { + return { + prompt_cache_key: prepared.promptCacheKey, + // Consumed by OpenAIProvider; it must never be forwarded as an API field. + // GPT-5.6+ uses it to end the cache before the changing query suffix. + _cascadeflow_knowledge_cache_prefix: prepared.systemPrefix, + }; + } + if (provider === 'anthropic') { + return { + cache_control: { + type: 'ephemeral', + ...(prepared.cacheTtl === '1h' ? { ttl: '1h' } : {}), + }, + }; + } + return {}; +} diff --git a/packages/core/src/providers/anthropic.ts b/packages/core/src/providers/anthropic.ts index 3c760340..7333001d 100644 --- a/packages/core/src/providers/anthropic.ts +++ b/packages/core/src/providers/anthropic.ts @@ -10,7 +10,7 @@ */ import { BaseProvider, type ProviderRequest, getHttpAgentOptions } from './base'; -import type { ProviderResponse, Tool, Message, ReasoningModelInfo } from '../types'; +import type { ProviderResponse, Tool, Message, ReasoningModelInfo, UsageDetails } from '../types'; import type { ModelConfig } from '../config'; import type { StreamChunk } from '../streaming'; @@ -26,28 +26,44 @@ try { } /** - * Anthropic pricing per 1M tokens (October 2025) - * Source: https://docs.claude.com/en/docs/about-claude/pricing + * Anthropic pricing per 1M tokens (August 2026) + * Source: https://platform.claude.com/docs/en/about-claude/pricing * - * Format: Blended rate (50% input, 50% output) + * Format: Input/output USD per million tokens. */ -const ANTHROPIC_PRICING: Record = { +const ANTHROPIC_PRICING: Record = { + // Claude 5 Series + 'claude-fable-5': { input: 10.0, output: 50.0 }, + 'claude-mythos-5': { input: 10.0, output: 50.0 }, + 'claude-opus-5': { input: 5.0, output: 25.0 }, + // Introductory price through 2026-08-31; invoices remain authoritative. + 'claude-sonnet-5': { input: 2.0, output: 10.0 }, + // Claude 4 Series - 'claude-opus-4.1': 45.0, // $15 in + $75 out = $45 blended - 'claude-opus-4': 45.0, // $15 in + $75 out = $45 blended - 'claude-sonnet-4.5': 9.0, // $3 in + $15 out = $9 blended - 'claude-sonnet-4': 9.0, // $3 in + $15 out = $9 blended + 'claude-opus-4-8': { input: 5.0, output: 25.0 }, + 'claude-opus-4-7': { input: 5.0, output: 25.0 }, + 'claude-opus-4-6': { input: 5.0, output: 25.0 }, + 'claude-opus-4-5': { input: 5.0, output: 25.0 }, + 'claude-opus-4.1': { input: 15.0, output: 75.0 }, + 'claude-opus-4': { input: 15.0, output: 75.0 }, + 'claude-sonnet-4-6': { input: 3.0, output: 15.0 }, + 'claude-sonnet-4-5': { input: 3.0, output: 15.0 }, + 'claude-sonnet-4.5': { input: 3.0, output: 15.0 }, + 'claude-sonnet-4': { input: 3.0, output: 15.0 }, + 'claude-haiku-4-5': { input: 1.0, output: 5.0 }, + 'claude-haiku-4.5': { input: 1.0, output: 5.0 }, // Claude 3.5 Series - 'claude-3-5-sonnet': 9.0, // $3 in + $15 out = $9 blended - 'claude-sonnet-3-5': 9.0, // Alternative naming - 'claude-3-5-haiku': 3.0, // $1 in + $5 out = $3 blended - 'claude-haiku-3-5': 3.0, // Alternative naming + 'claude-3-7-sonnet': { input: 3.0, output: 15.0 }, + 'claude-3-5-sonnet': { input: 3.0, output: 15.0 }, + 'claude-sonnet-3-5': { input: 3.0, output: 15.0 }, + 'claude-3-5-haiku': { input: 0.8, output: 4.0 }, + 'claude-haiku-3-5': { input: 0.8, output: 4.0 }, // Claude 3 Series - 'claude-3-opus': 45.0, // $15 in + $75 out = $45 blended - 'claude-3-sonnet': 9.0, // $3 in + $15 out = $9 blended - 'claude-3-haiku': 0.75, // $0.25 in + $1.25 out = $0.75 blended + 'claude-3-opus': { input: 15.0, output: 75.0 }, + 'claude-3-sonnet': { input: 3.0, output: 15.0 }, + 'claude-3-haiku': { input: 0.25, output: 1.25 }, }; /** @@ -353,6 +369,12 @@ export class AnthropicProvider extends BaseProvider { prompt_tokens: completion.usage.input_tokens, completion_tokens: completion.usage.output_tokens, total_tokens: completion.usage.input_tokens + completion.usage.output_tokens, + cached_input_tokens: completion.usage.cache_read_input_tokens, + cache_write_input_tokens: completion.usage.cache_creation_input_tokens, + cache_write_5m_input_tokens: + (completion.usage as any).cache_creation?.ephemeral_5m_input_tokens, + cache_write_1h_input_tokens: + (completion.usage as any).cache_creation?.ephemeral_1h_input_tokens, }, finish_reason: completion.stop_reason || undefined, tool_calls: toolCalls.length > 0 ? toolCalls : undefined, @@ -438,6 +460,12 @@ export class AnthropicProvider extends BaseProvider { prompt_tokens: completion.usage.input_tokens, completion_tokens: completion.usage.output_tokens, total_tokens: completion.usage.input_tokens + completion.usage.output_tokens, + cached_input_tokens: completion.usage.cache_read_input_tokens, + cache_write_input_tokens: completion.usage.cache_creation_input_tokens, + cache_write_5m_input_tokens: + completion.usage.cache_creation?.ephemeral_5m_input_tokens, + cache_write_1h_input_tokens: + completion.usage.cache_creation?.ephemeral_1h_input_tokens, }, finish_reason: completion.stop_reason || undefined, tool_calls: toolCalls.length > 0 ? toolCalls : undefined, @@ -449,18 +477,50 @@ export class AnthropicProvider extends BaseProvider { } calculateCost(promptTokens: number, completionTokens: number, model: string): number { - const totalTokens = promptTokens + completionTokens; - const modelLower = model.toLowerCase(); + const pricing = this.getModelPricing(model); + return (promptTokens * pricing.input + completionTokens * pricing.output) / 1_000_000; + } - // Find matching rate (try prefix match) - for (const [modelPrefix, rate] of Object.entries(ANTHROPIC_PRICING)) { - if (modelLower.includes(modelPrefix)) { - return (totalTokens / 1_000_000) * rate; - } + calculateCostFromUsage(usage: UsageDetails, model: string): number { + const pricing = this.getModelPricing(model); + const promptTokens = Math.max(usage.prompt_tokens ?? usage.input_tokens ?? 0, 0); + const completionTokens = Math.max( + usage.completion_tokens ?? usage.output_tokens ?? 0, + 0 + ); + const readTokens = Math.max(usage.cached_input_tokens ?? 0, 0); + const writeTokens = Math.max(usage.cache_write_input_tokens ?? 0, 0); + const write5m = Math.min(Math.max(usage.cache_write_5m_input_tokens ?? 0, 0), writeTokens); + const write1h = Math.min( + Math.max(usage.cache_write_1h_input_tokens ?? 0, 0), + Math.max(writeTokens - write5m, 0) + ); + + if (readTokens === 0 && writeTokens === 0) { + return this.calculateCost(promptTokens, completionTokens, model); } - // Default to Sonnet pricing if unknown - return (totalTokens / 1_000_000) * 9.0; + const unclassifiedWrites = Math.max(writeTokens - write5m - write1h, 0); + const billedInputTokens = + promptTokens + + readTokens * 0.1 + + write5m * 1.25 + + write1h * 2.0 + + unclassifiedWrites * 1.25; + return (billedInputTokens * pricing.input + completionTokens * pricing.output) / 1_000_000; + } + + private getModelPricing(model: string): { input: number; output: number } { + const modelLower = model.toLowerCase().split('/').pop() ?? model.toLowerCase(); + let longestMatch = ''; + let pricing: { input: number; output: number } | undefined; + for (const [prefix, value] of Object.entries(ANTHROPIC_PRICING)) { + if (modelLower.startsWith(prefix) && prefix.length > longestMatch.length) { + longestMatch = prefix; + pricing = value; + } + } + return pricing ?? { input: 3.0, output: 15.0 }; } /** diff --git a/packages/core/src/providers/base.ts b/packages/core/src/providers/base.ts index 94cd7906..4aca41c2 100644 --- a/packages/core/src/providers/base.ts +++ b/packages/core/src/providers/base.ts @@ -2,7 +2,7 @@ * Base provider interface and utilities */ -import type { Message, Tool, ProviderResponse, HttpConfig } from '../types'; +import type { Message, Tool, ProviderResponse, HttpConfig, UsageDetails } from '../types'; import type { ModelConfig } from '../config'; import type { StreamChunk } from '../streaming'; @@ -182,6 +182,9 @@ export interface Provider { */ calculateCost(promptTokens: number, completionTokens: number, model: string): number; + /** Calculate actual cost from provider-reported cache token categories when available. */ + calculateCostFromUsage?(usage: UsageDetails, model: string): number; + /** * Check if provider is available (API key set, etc.) */ @@ -201,6 +204,10 @@ export abstract class BaseProvider implements Provider { abstract generate(request: ProviderRequest): Promise; abstract calculateCost(promptTokens: number, completionTokens: number, model: string): number; + calculateCostFromUsage(usage: UsageDetails, model: string): number { + return this.calculateCost(usage.prompt_tokens, usage.completion_tokens, model); + } + isAvailable(): boolean { return !!this.config.apiKey || !!process.env[`${this.name.toUpperCase()}_API_KEY`]; } diff --git a/packages/core/src/providers/openai.ts b/packages/core/src/providers/openai.ts index 63d5920f..c28443e8 100644 --- a/packages/core/src/providers/openai.ts +++ b/packages/core/src/providers/openai.ts @@ -6,7 +6,7 @@ */ import { BaseProvider, type ProviderRequest, getHttpAgentOptions } from './base'; -import type { ProviderResponse, Tool, Message, ReasoningModelInfo } from '../types'; +import type { ProviderResponse, Tool, Message, ReasoningModelInfo, UsageDetails } from '../types'; import type { ModelConfig } from '../config'; import type { StreamChunk } from '../streaming'; @@ -27,11 +27,64 @@ try { type ChatCompletionMessageParam = any; // Simplified for MVP type ChatCompletionTool = any; // Simplified for MVP +const KNOWLEDGE_CACHE_PREFIX_FIELD = '_cascadeflow_knowledge_cache_prefix'; + +function usesExplicitPromptCacheBreakpoints(model: string): boolean { + const modelName = model.split('/').pop() ?? model; + const match = /^gpt-(\d+)(?:\.(\d+))?/i.exec(modelName); + if (!match) return false; + const major = Number(match[1]); + const minor = Number(match[2] ?? 0); + return major > 5 || (major === 5 && minor >= 6); +} + +function prepareOpenAICacheExtra( + model: string, + extra: Record | undefined, + messages: ChatCompletionMessageParam[] +): Record { + const providerExtra = { ...(extra ?? {}) }; + const stablePrefix = providerExtra[KNOWLEDGE_CACHE_PREFIX_FIELD]; + delete providerExtra[KNOWLEDGE_CACHE_PREFIX_FIELD]; + + if (typeof stablePrefix !== 'string' || !usesExplicitPromptCacheBreakpoints(model)) { + return providerExtra; + } + + for (const message of messages) { + if (typeof message.content !== 'string') continue; + const start = message.content.indexOf(stablePrefix); + if (start < 0) continue; + const splitAt = start + stablePrefix.length; + message.content = [ + { + type: 'text', + text: message.content.slice(0, splitAt), + prompt_cache_breakpoint: { mode: 'explicit' }, + }, + ...(message.content.slice(splitAt) + ? [{ type: 'text', text: message.content.slice(splitAt) }] + : []), + ]; + // Avoid a billable implicit write for the changing latest user message. + providerExtra.prompt_cache_options = { mode: 'explicit' }; + break; + } + + return providerExtra; +} + /** * OpenAI pricing per 1K tokens (as of January 2025) * Source: https://openai.com/api/pricing/ */ const OPENAI_PRICING: Record = { + // GPT-5.6 standard short-context pricing (August 2026) + 'gpt-5.6-sol': { input: 0.005, output: 0.030 }, + 'gpt-5.6-terra': { input: 0.002, output: 0.012 }, + 'gpt-5.6-luna': { input: 0.0002, output: 0.0012 }, + 'gpt-5.6': { input: 0.005, output: 0.030 }, + // GPT-5 series (current flagship - released August 2025) // 50% cheaper input than GPT-4o, superior performance on coding, reasoning, math 'gpt-5': { input: 0.00125, output: 0.010 }, @@ -219,6 +272,7 @@ export class OpenAIProvider extends BaseProvider { const tools = request.tools ? this.convertTools(request.tools) : undefined; const modelName = request.model || this.config.name; + const providerExtra = prepareOpenAICacheExtra(modelName, request.extra, chatMessages); // GPT-5 series doesn't support logprobs yet (as of January 2025) const supportsLogprobs = !modelName.startsWith('gpt-5'); const isGpt5 = modelName.startsWith('gpt-5'); @@ -229,7 +283,7 @@ export class OpenAIProvider extends BaseProvider { messages: chatMessages, tools, stream: true, - ...request.extra, + ...providerExtra, }; // GPT-5 only supports temperature=1 (default), doesn't allow custom values @@ -321,6 +375,7 @@ export class OpenAIProvider extends BaseProvider { const tools = request.tools ? this.convertTools(request.tools) : undefined; const modelName = request.model || this.config.name; + const providerExtra = prepareOpenAICacheExtra(modelName, request.extra, chatMessages); // GPT-5 series doesn't support logprobs yet (as of January 2025) const supportsLogprobs = !modelName.startsWith('gpt-5'); const isGpt5 = modelName.startsWith('gpt-5'); @@ -331,7 +386,7 @@ export class OpenAIProvider extends BaseProvider { messages: chatMessages, tools, stream: true, - ...request.extra, + ...providerExtra, }; // GPT-5 only supports temperature=1 (default), doesn't allow custom values @@ -473,6 +528,7 @@ export class OpenAIProvider extends BaseProvider { messages, modelInfo.supportsSystemMessages ? request.systemPrompt : undefined ); + const providerExtra = prepareOpenAICacheExtra(modelName, request.extra, chatMessages); // If system prompt provided but not supported, prepend to first user message if (!modelInfo.supportsSystemMessages && request.systemPrompt) { @@ -498,7 +554,7 @@ export class OpenAIProvider extends BaseProvider { model: modelName, messages: chatMessages, tools, - ...request.extra, + ...providerExtra, }; // GPT-5 only supports temperature=1 (default), doesn't allow custom values @@ -544,6 +600,10 @@ export class OpenAIProvider extends BaseProvider { prompt_tokens: completion.usage.prompt_tokens, completion_tokens: completion.usage.completion_tokens, total_tokens: completion.usage.total_tokens, + cached_input_tokens: completion.usage.prompt_tokens_details?.cached_tokens, + cache_write_input_tokens: + completion.usage.cache_write_tokens ?? + (completion.usage.prompt_tokens_details as any)?.cache_write_tokens, reasoning_tokens: completion.usage.completion_tokens_details?.reasoning_tokens, completion_tokens_details: completion.usage.completion_tokens_details, } @@ -580,6 +640,7 @@ export class OpenAIProvider extends BaseProvider { const tools = request.tools ? this.convertTools(request.tools) : undefined; const modelName = request.model || this.config.name; + const providerExtra = prepareOpenAICacheExtra(modelName, request.extra, chatMessages); // GPT-5 series doesn't support logprobs yet (as of January 2025) const supportsLogprobs = !modelName.startsWith('gpt-5'); const isGpt5 = modelName.startsWith('gpt-5'); @@ -589,7 +650,7 @@ export class OpenAIProvider extends BaseProvider { model: modelName, messages: chatMessages, tools, - ...request.extra, + ...providerExtra, }; // GPT-5 only supports temperature=1 (default), doesn't allow custom values @@ -645,6 +706,10 @@ export class OpenAIProvider extends BaseProvider { prompt_tokens: completion.usage.prompt_tokens, completion_tokens: completion.usage.completion_tokens, total_tokens: completion.usage.total_tokens, + cached_input_tokens: completion.usage.prompt_tokens_details?.cached_tokens, + cache_write_input_tokens: + completion.usage.cache_write_tokens ?? + completion.usage.prompt_tokens_details?.cache_write_tokens, } : undefined, finish_reason: choice.finish_reason, @@ -679,34 +744,7 @@ export class OpenAIProvider extends BaseProvider { model: string, _reasoningTokens?: number ): number { - // Normalize model name to lowercase for case-insensitive matching - const modelLower = model.toLowerCase(); - - // Find model-specific pricing (exact match, case-insensitive) - let pricing = OPENAI_PRICING[modelLower]; - - // Try prefix matching for versioned models - // Match the LONGEST prefix to avoid "gpt-4o" matching "gpt-4o-mini" - if (!pricing) { - let longestMatch = ''; - let longestMatchPricing = null; - - for (const [key, value] of Object.entries(OPENAI_PRICING)) { - if (modelLower.startsWith(key) && key.length > longestMatch.length) { - longestMatch = key; - longestMatchPricing = value; - } - } - - if (longestMatchPricing) { - pricing = longestMatchPricing; - } - } - - // Fallback to gpt-4o-mini for unknown models - if (!pricing) { - pricing = OPENAI_PRICING['gpt-4o-mini']; - } + const pricing = this.getModelPricing(model); // Note: For o1/o3 models, reasoning tokens are already included in completion_tokens // from the API, so we don't need to add them separately. The API returns: @@ -718,6 +756,48 @@ export class OpenAIProvider extends BaseProvider { return totalCost; } + calculateCostFromUsage(usage: UsageDetails, model: string): number { + const pricing = this.getModelPricing(model); + const promptTokens = Math.max(usage.prompt_tokens ?? usage.input_tokens ?? 0, 0); + const completionTokens = Math.max( + usage.completion_tokens ?? usage.output_tokens ?? 0, + 0 + ); + const cachedTokens = Math.min(Math.max(usage.cached_input_tokens ?? 0, 0), promptTokens); + const writeTokens = Math.min( + Math.max(usage.cache_write_input_tokens ?? 0, 0), + Math.max(promptTokens - cachedTokens, 0) + ); + + if (cachedTokens === 0 && writeTokens === 0) { + return this.calculateCost(promptTokens, completionTokens, model); + } + + const uncachedTokens = Math.max(promptTokens - cachedTokens - writeTokens, 0); + const writeMultiplier = usesExplicitPromptCacheBreakpoints(model) ? 1.25 : 1.0; + const inputCost = + ((uncachedTokens + cachedTokens * 0.1 + writeTokens * writeMultiplier) / 1000) * + pricing.input; + const outputCost = (completionTokens / 1000) * pricing.output; + return inputCost + outputCost; + } + + private getModelPricing(model: string): { input: number; output: number } { + const modelLower = model.toLowerCase().split('/').pop() ?? model.toLowerCase(); + const exact = OPENAI_PRICING[modelLower]; + if (exact) return exact; + + let longestMatch = ''; + let pricing: { input: number; output: number } | undefined; + for (const [key, value] of Object.entries(OPENAI_PRICING)) { + if (modelLower.startsWith(key) && key.length > longestMatch.length) { + longestMatch = key; + pricing = value; + } + } + return pricing ?? OPENAI_PRICING['gpt-4o-mini']; + } + /** * Convert generic messages to OpenAI chat format */ diff --git a/packages/core/src/streaming.ts b/packages/core/src/streaming.ts index b1bae3ee..ea43a051 100644 --- a/packages/core/src/streaming.ts +++ b/packages/core/src/streaming.ts @@ -1,3 +1,5 @@ +import type { KnowledgeSnapshot } from './knowledge-cache'; + /** * Streaming types and interfaces for cascadeflow * @@ -151,6 +153,9 @@ export interface StreamOptions { /** System prompt */ systemPrompt?: string; + /** Immutable request-scoped knowledge used by every routed model. */ + knowledge?: string | KnowledgeSnapshot; + /** Tools available */ tools?: any[]; diff --git a/packages/core/src/types.ts b/packages/core/src/types.ts index a8ff7e04..7868d465 100644 --- a/packages/core/src/types.ts +++ b/packages/core/src/types.ts @@ -110,6 +110,7 @@ export interface Usage { input_tokens: number; output_tokens: number; cached_input_tokens: number; + cache_write_input_tokens: number; total_tokens: number; } @@ -118,10 +119,13 @@ export function toCanonicalUsage(usage?: Partial | Partial) const outputTokens = (usage as any)?.output_tokens ?? (usage as any)?.completion_tokens ?? 0; const cachedInputTokens = (usage as any)?.cached_input_tokens ?? (usage as any)?.cache_read_input_tokens ?? 0; + const cacheWriteInputTokens = + (usage as any)?.cache_write_input_tokens ?? (usage as any)?.cache_creation_input_tokens ?? 0; return { input_tokens: inputTokens, output_tokens: outputTokens, cached_input_tokens: cachedInputTokens, + cache_write_input_tokens: cacheWriteInputTokens, total_tokens: inputTokens + outputTokens, }; } @@ -136,6 +140,9 @@ export interface UsageDetails { input_tokens?: number; output_tokens?: number; cached_input_tokens?: number; + cache_write_input_tokens?: number; + cache_write_5m_input_tokens?: number; + cache_write_1h_input_tokens?: number; reasoning_tokens?: number; // For OpenAI o1/o3 models completion_tokens_details?: { reasoning_tokens?: number; diff --git a/pyproject.toml b/pyproject.toml index 7968f0ce..2b25cb8f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -121,6 +121,9 @@ google-adk = ["google-adk>=1.0.0; python_version >= '3.10'"] # PydanticAI integration (opt-in) pydantic-ai = ["pydantic-ai>=0.1.0; python_version >= '3.10'"] +# Model Context Protocol server integration (opt-in; v1 FastMCP API) +mcp = ["mcp>=1.27,<2; python_version >= '3.10'"] + # Development tools (includes rich for terminal output) dev = [ "pytest>=7.4.0", diff --git a/tests/test_anthropic.py b/tests/test_anthropic.py index 46c87899..156f8d0c 100644 --- a/tests/test_anthropic.py +++ b/tests/test_anthropic.py @@ -139,6 +139,24 @@ def test_estimate_cost_haiku(self, anthropic_provider): # Uses blended pricing assert 0.0005 < cost < 0.0010 # Approximately $0.00075/1K tokens + def test_cache_cost_distinguishes_reads_and_one_hour_writes(self, anthropic_provider): + cost = anthropic_provider._calculate_response_cost( + model="claude-sonnet-4-5-20250929", + prompt_tokens=100, + completion_tokens=10, + usage={ + "input_tokens": 100, + "output_tokens": 10, + "cache_read_input_tokens": 1000, + "cache_creation_input_tokens": 1000, + "cache_creation": {"ephemeral_1h_input_tokens": 1000}, + }, + ) + + # Input: (100 uncached + 1000 reads at 0.1x + 1000 writes at 2x) * $3/M. + # Output: 10 * $15/M. + assert cost == pytest.approx(0.00675) + if __name__ == "__main__": pytest.main([__file__, "-v"]) diff --git a/tests/test_knowledge_context.py b/tests/test_knowledge_context.py new file mode 100644 index 00000000..baa9c129 --- /dev/null +++ b/tests/test_knowledge_context.py @@ -0,0 +1,137 @@ +from __future__ import annotations + +from unittest.mock import patch + +import pytest + +from cascadeflow.agent import CascadeAgent +from cascadeflow.context import KnowledgeCache, KnowledgeSnapshot, provider_cache_kwargs +from cascadeflow.providers.base import ModelResponse +from cascadeflow.schema.config import ModelConfig + + +def test_snapshot_is_content_addressed_and_reused() -> None: + cache = KnowledgeCache(max_entries=2) + first, first_hit = cache.prepare(KnowledgeSnapshot("alpha", key="docs")) + second, second_hit = cache.prepare(KnowledgeSnapshot("alpha", key="docs")) + + assert not first_hit + assert second_hit + assert first is second + assert cache.get_stats() == {"size": 1, "hits": 1, "misses": 1, "evictions": 0} + + +def test_explicit_version_cannot_silently_change_content() -> None: + cache = KnowledgeCache() + cache.prepare(KnowledgeSnapshot("current", key="docs", version="v1")) + + with pytest.raises(ValueError, match="version collision"): + cache.prepare(KnowledgeSnapshot("stale-or-different", key="docs", version="v1")) + + +def test_provider_key_remains_content_bound_after_local_eviction() -> None: + cache = KnowledgeCache(max_entries=1) + old, _ = cache.prepare(KnowledgeSnapshot("old", key="docs", version="v1")) + cache.prepare(KnowledgeSnapshot("other", key="other", version="v1")) + new, _ = cache.prepare(KnowledgeSnapshot("new", key="docs", version="v1")) + + assert old.prompt_cache_key != new.prompt_cache_key + + +def test_provider_cache_hints_are_capability_specific() -> None: + prepared, _ = KnowledgeCache().prepare( + KnowledgeSnapshot("shared", key="docs", version="v3", cache_ttl="1h") + ) + + assert provider_cache_kwargs("openai", prepared) == { + "prompt_cache_key": prepared.prompt_cache_key, + "_cascadeflow_knowledge_cache_prefix": prepared.system_prefix, + } + assert provider_cache_kwargs("anthropic", prepared) == { + "cache_control": {"type": "ephemeral", "ttl": "1h"} + } + assert provider_cache_kwargs("ollama", prepared) == {} + + +@pytest.mark.asyncio +async def test_agent_switches_knowledge_without_leaking_previous_snapshot() -> None: + model = ModelConfig(name="gpt-test", provider="openai", cost=0.001) + + class RecordingProvider: + def __init__(self) -> None: + self.calls: list[dict] = [] + + async def complete(self, **kwargs): + self.calls.append(kwargs) + return ModelResponse( + content="ok", + model="gpt-test", + provider="openai", + cost=0.001, + tokens_used=2, + confidence=0.9, + metadata={"input_tokens": 1, "output_tokens": 1}, + ) + + provider = RecordingProvider() + with patch("cascadeflow.agent.PROVIDER_REGISTRY") as registry: + registry.__getitem__.return_value = lambda: provider + registry.__contains__.return_value = True + agent = CascadeAgent(models=[model], enable_cascade=False) + agent.providers = {"openai": provider} + agent.model_providers = {model.name: provider} + + alpha = KnowledgeSnapshot("ALPHA_ONLY", key="tenant", version="a") + beta = KnowledgeSnapshot("BETA_ONLY", key="tenant", version="b") + + first = await agent.run("question", force_direct=True, knowledge=alpha) + second = await agent.run("question", force_direct=True, knowledge=beta) + third = await agent.run("question", force_direct=True, knowledge=beta) + + first_prompt = provider.calls[0]["prompt"] + second_prompt = provider.calls[1]["prompt"] + assert "ALPHA_ONLY" in first_prompt + assert "BETA_ONLY" not in first_prompt + assert "BETA_ONLY" in second_prompt + assert "ALPHA_ONLY" not in second_prompt + assert provider.calls[0]["prompt_cache_key"] != provider.calls[1]["prompt_cache_key"] + assert first.metadata["knowledge"]["local_cache_hit"] is False + assert second.metadata["knowledge"]["local_cache_hit"] is False + assert third.metadata["knowledge"]["local_cache_hit"] is True + + +@pytest.mark.asyncio +async def test_system_prompt_is_kept_after_stable_knowledge_prefix() -> None: + model = ModelConfig(name="local-test", provider="ollama", cost=0.0) + + class RecordingProvider: + async def complete(self, **kwargs): + self.kwargs = kwargs + return ModelResponse( + content="ok", + model="local-test", + provider="ollama", + cost=0.0, + tokens_used=2, + confidence=0.9, + ) + + provider = RecordingProvider() + with patch("cascadeflow.agent.PROVIDER_REGISTRY") as registry: + registry.__getitem__.return_value = lambda: provider + registry.__contains__.return_value = True + agent = CascadeAgent(models=[model], enable_cascade=False) + agent.providers = {"ollama": provider} + agent.model_providers = {model.name: provider} + + await agent.run( + "question", + force_direct=True, + system_prompt="SYSTEM_RULE", + knowledge=KnowledgeSnapshot("KNOWLEDGE", key="docs", version="1"), + ) + + prompt = provider.kwargs["prompt"] + assert prompt.index("KNOWLEDGE") < prompt.index("SYSTEM_RULE") < prompt.index("question") + assert "prompt_cache_key" not in provider.kwargs + assert "cache_control" not in provider.kwargs diff --git a/tests/test_mcp_integration.py b/tests/test_mcp_integration.py new file mode 100644 index 00000000..d9bd1993 --- /dev/null +++ b/tests/test_mcp_integration.py @@ -0,0 +1,103 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any +from unittest.mock import patch + +import pytest + +from cascadeflow.context import KnowledgeSnapshot +from cascadeflow.integrations.mcp import create_mcp_server + + +class FakeFastMCP: + def __init__(self, name: str, **kwargs: Any) -> None: + self.name = name + self.kwargs = kwargs + self.tools: dict[str, Any] = {} + + def tool(self, **kwargs: Any): + def decorate(function): + self.tools[function.__name__] = function + return function + + return decorate + + +@dataclass +class FakeResult: + content: str = "answer" + model_used: str = "draft" + cascaded: bool = False + draft_accepted: bool = True + routing_strategy: str = "cascade" + total_cost: float = 0.001 + cost_saved: float = 0.009 + latency_ms: float = 12.0 + metadata: dict[str, Any] = field(default_factory=lambda: {"knowledge": {"identity": "docs:v2"}}) + + +class FakeAgent: + def __init__(self) -> None: + self.calls: list[dict[str, Any]] = [] + + async def run(self, query: str, **kwargs: Any) -> FakeResult: + self.calls.append({"query": query, **kwargs}) + return FakeResult() + + +@pytest.mark.asyncio +async def test_mcp_tool_resolves_knowledge_server_side_and_returns_compact_trace() -> None: + agent = FakeAgent() + + async def resolve(key: str, version: str | None) -> KnowledgeSnapshot: + return KnowledgeSnapshot("private docs", key=key, version=version) + + with patch("cascadeflow.integrations.mcp._load_fastmcp", return_value=FakeFastMCP): + server = create_mcp_server(agent, knowledge_resolver=resolve) + + response = await server.tools["cascadeflow_run"]( + "current question", + knowledge_key="docs", + knowledge_version="v2", + conversation_context="Only the relevant prior fact.", + ) + + call = agent.calls[0] + assert call["knowledge"].content == "private docs" + assert call["knowledge"].identity == "docs:v2" + assert call["messages"][-1] == {"role": "user", "content": "current question"} + assert response["content"] == "answer" + assert response["knowledge"] == {"identity": "docs:v2"} + + +@pytest.mark.asyncio +async def test_mcp_tool_requires_resolver_and_bounds_conversation_handoff() -> None: + with patch("cascadeflow.integrations.mcp._load_fastmcp", return_value=FakeFastMCP): + server = create_mcp_server(FakeAgent(), max_context_chars=5) + + tool = server.tools["cascadeflow_run"] + with pytest.raises(ValueError, match="knowledge_resolver"): + await tool("query", knowledge_key="docs") + with pytest.raises(ValueError, match="concise relevant handoff"): + await tool("query", conversation_context="too long") + + +@pytest.mark.asyncio +async def test_mcp_sdk_discovers_and_calls_cascadeflow_tool() -> None: + pytest.importorskip("mcp") + + server = create_mcp_server(FakeAgent()) + tools = await server.list_tools() + + tool = next(item for item in tools if item.name == "cascadeflow_run") + assert tool.title == "Run cascadeflow" + assert tool.outputSchema["type"] == "object" + assert tool.annotations.readOnlyHint is True + assert tool.annotations.destructiveHint is False + assert tool.annotations.openWorldHint is True + + content, structured = await server.call_tool("cascadeflow_run", {"query": "current question"}) + assert content[0].type == "text" + assert structured["content"] == "answer" + assert structured["model_used"] == "draft" diff --git a/tests/test_openai.py b/tests/test_openai.py index 3975c452..4cb8a769 100644 --- a/tests/test_openai.py +++ b/tests/test_openai.py @@ -92,6 +92,106 @@ async def test_complete_with_system_prompt(self, openai_provider, mock_openai_re assert messages[0]["role"] == "system" assert messages[1]["role"] == "user" + @pytest.mark.asyncio + async def test_gpt56_marks_only_stable_knowledge_for_explicit_cache(self, openai_provider): + """GPT-5.6 must not create billable cache writes for each changing query.""" + stable = '\nKNOWN\n' + response_data = { + "status": "completed", + "model": "gpt-5.6-terra", + "output": [ + { + "type": "message", + "content": [{"type": "output_text", "text": "ok"}], + } + ], + "usage": { + "input_tokens": 1100, + "output_tokens": 1, + "total_tokens": 1101, + "input_tokens_details": { + "cached_tokens": 1024, + "cache_write_tokens": 0, + }, + }, + } + with patch.object(openai_provider.client, "post") as mock_post: + mock_response = MagicMock() + mock_response.json.return_value = response_data + mock_response.raise_for_status = MagicMock() + mock_post.return_value = mock_response + + result = await openai_provider._complete_impl( + prompt=f"System: {stable}\nUser: changing question", + model="gpt-5.6-terra", + prompt_cache_key="knowledge-key", + _cascadeflow_knowledge_cache_prefix=stable, + ) + + payload = mock_post.call_args.kwargs["json"] + blocks = payload["input"][0]["content"] + assert payload["prompt_cache_options"] == {"mode": "explicit"} + assert blocks[0]["prompt_cache_breakpoint"] == {"mode": "explicit"} + assert blocks[0]["text"].endswith(stable) + assert blocks[1]["text"] == "\nUser: changing question" + assert "_cascadeflow_knowledge_cache_prefix" not in str(payload) + assert result.metadata["cached_input_tokens"] == 1024 + assert result.metadata["cache_write_input_tokens"] == 0 + # 76 uncached + 1024 cached at 10%, plus one output token. + assert result.cost == pytest.approx(0.0003688) + + def test_gpt56_cache_write_cost_uses_provider_token_category(self, openai_provider): + cost = openai_provider._calculate_response_cost( + model="gpt-5.6-terra", + prompt_tokens=1100, + completion_tokens=1, + usage={ + "input_tokens": 1100, + "output_tokens": 1, + "input_tokens_details": { + "cached_tokens": 0, + "cache_write_tokens": 1024, + }, + }, + ) + + # 76 uncached + 1024 writes at 1.25x, plus one output token. + assert cost == pytest.approx(0.002724) + + @pytest.mark.asyncio + async def test_pre_gpt56_keeps_automatic_cache_shape(self, openai_provider): + """Older OpenAI models reject GPT-5.6-only breakpoint fields.""" + stable = '\nKNOWN\n' + response_data = { + "status": "completed", + "model": "gpt-5.5", + "output": [ + { + "type": "message", + "content": [{"type": "output_text", "text": "ok"}], + } + ], + "usage": {"input_tokens": 1100, "output_tokens": 1, "total_tokens": 1101}, + } + with patch.object(openai_provider.client, "post") as mock_post: + mock_response = MagicMock() + mock_response.json.return_value = response_data + mock_response.raise_for_status = MagicMock() + mock_post.return_value = mock_response + + await openai_provider._complete_impl( + prompt=f"System: {stable}\nUser: changing question", + model="gpt-5.5", + prompt_cache_key="knowledge-key", + _cascadeflow_knowledge_cache_prefix=stable, + ) + + payload = mock_post.call_args.kwargs["json"] + assert payload["prompt_cache_key"] == "knowledge-key" + assert "prompt_cache_options" not in payload + assert isinstance(payload["input"][0]["content"], str) + assert "_cascadeflow_knowledge_cache_prefix" not in str(payload) + @pytest.mark.asyncio async def test_complete_http_error(self, openai_provider): """Test handling of HTTP errors.""" diff --git a/tests/test_pricing_resolver.py b/tests/test_pricing_resolver.py index 7c73a7cf..fe0f0cd9 100644 --- a/tests/test_pricing_resolver.py +++ b/tests/test_pricing_resolver.py @@ -35,8 +35,10 @@ def test_usage_from_payload_maps_legacy_fields(): "prompt_tokens": 12, "completion_tokens": 8, "cache_read_input_tokens": 3, + "cache_creation_input_tokens": 5, } ) assert usage.input_tokens == 12 assert usage.output_tokens == 8 assert usage.cached_input_tokens == 3 + assert usage.cache_write_input_tokens == 5 diff --git a/tests/test_reasoning_models.py b/tests/test_reasoning_models.py index 6b425337..7d0f0641 100644 --- a/tests/test_reasoning_models.py +++ b/tests/test_reasoning_models.py @@ -329,10 +329,10 @@ def test_claude_3_5_sonnet_cost(self): def test_claude_3_5_haiku_cost(self): """Test claude-3-5-haiku cost calculation.""" provider = AnthropicProvider(api_key="test") - # Blended rate: $3.00 per 1M tokens - # 2M tokens total = $6.00 + # Blended rate: $2.40 per 1M tokens ($0.80 input / $4 output) + # 2M tokens total = $4.80 cost = provider.estimate_cost(tokens=2000000, model="claude-3-5-haiku-20241022") - assert abs(cost - 6.0) < 0.001 + assert abs(cost - 4.8) < 0.001 def test_prefix_matching_versioned_models(self): """Test prefix matching for versioned Claude 3.5 models.""" @@ -375,7 +375,7 @@ class TestAnthropicPricingMatrix: ("claude-sonnet-4", 9.0), # Claude 3.5 Series ("claude-3-5-sonnet", 9.0), - ("claude-3-5-haiku", 3.0), + ("claude-3-5-haiku", 2.4), # Claude 3 Series ("claude-3-opus", 45.0), ("claude-3-sonnet", 9.0),