From cb411eaadfbade4cf78356f6bac98e1771779d61 Mon Sep 17 00:00:00 2001 From: John Toman Date: Tue, 21 Jul 2026 12:10:55 -0700 Subject: [PATCH 1/4] LLM backend plugins Instead of hardcoding "anthropic" and "openai", we generalize ProviderKind into an abstract class, one with two key dynamic dispatches that encapsulate the behavior we would case split over the two constants for: selecting the correct memory tool implementation and text/file document rendering. This new ABC is called ProviderServices, and is threaded through everywhere the old ProviderKind was. Further, the actual providers are now auto-discovered and loaded via importlib.metadata entrypoints under `certora.autoprove.llm_provider`. These entrypoints are expected to resolve to an instance of ProviderSpec, which describes how to select which provider to use for a model name, and the provider itself. The provider is responsible for taking the model name and configuration options and producing a modelprovider; itself encapsulates the model "builder" (itself parameterized over caching/thinking behavior) and the "provider services", which is the ABC described above. As a validation of this approach, the existing OpenAI and anthropic providers have been moved to be "built in plugins". In principle adding Kimi, Qwen, Deepseek, etc. should simply be a matter of implementing the various Provider apis. Common base classes/abstractions are provided for the anticipated common cases to cut down on boilerplate. While doing this refactoring, took the opportunity to clean up some apis; including removing the (deprecated and very broken) get_checkpointer, the composer.workflow.providers shim, and other now unused and orphan llm builders. --- composer/cli/console_codegen.py | 4 +- composer/cli/tui_codegen.py | 4 +- composer/cli/tui_pipeline.py | 7 +- composer/input/files.py | 63 ++++--------- composer/llm/anthropic.py | 65 +++++++++++-- composer/llm/api.py | 7 ++ composer/llm/openai.py | 45 +++++++-- composer/llm/parsing.py | 25 +++++ composer/llm/provider.py | 75 ++++++++------- composer/llm/registry.py | 100 +++++--------------- composer/natreq/extractor.py | 2 - composer/pipeline/cli.py | 2 +- composer/spec/services.py | 1 - composer/spec/source/author.py | 2 +- composer/testing/harness_tape.py | 65 +++++++++---- composer/workflow/provider.py | 11 --- composer/workflow/services.py | 153 ++----------------------------- pyproject.toml | 6 ++ sanity_analyzer/analysis.py | 21 ++--- 19 files changed, 290 insertions(+), 368 deletions(-) create mode 100644 composer/llm/api.py create mode 100644 composer/llm/parsing.py delete mode 100644 composer/workflow/provider.py diff --git a/composer/cli/console_codegen.py b/composer/cli/console_codegen.py index 43c54c02..0be36166 100644 --- a/composer/cli/console_codegen.py +++ b/composer/cli/console_codegen.py @@ -11,7 +11,7 @@ import sys from composer.input.parsing import fresh_workflow_argument_parser, upload_input -from composer.llm.registry import get_provider_for, uploader_for +from composer.llm.registry import get_provider_for from composer.workflow.executor import execute_ai_composer_workflow from composer.workflow.types import WorkflowSuccess from composer.ui.console import ConsoleHandler @@ -23,7 +23,7 @@ async def _main() -> int: args = parser.parse_args() llm = get_provider_for(options=args) - input_data = await upload_input(uploader_for(llm.provider), args) + input_data = await upload_input(llm.provider.uploader(), args) handler = ConsoleHandler(capture_prover_output=args.prover_capture_output) with tool_context(): diff --git a/composer/cli/tui_codegen.py b/composer/cli/tui_codegen.py index c0ba436c..a176c53d 100644 --- a/composer/cli/tui_codegen.py +++ b/composer/cli/tui_codegen.py @@ -6,7 +6,7 @@ import sys from composer.input.parsing import fresh_workflow_argument_parser, upload_input -from composer.llm.registry import get_provider_for, uploader_for +from composer.llm.registry import get_provider_for from composer.workflow.executor import execute_ai_composer_workflow from composer.ui.codegen_rich import CodeGenRichApp from composer.ui.tool_display import tool_context @@ -23,7 +23,7 @@ async def _main() -> int: llm = get_provider_for(options=args) - uploader = uploader_for(llm.provider) + uploader = llm.provider.uploader() input_data = await upload_input(uploader, args) diff --git a/composer/cli/tui_pipeline.py b/composer/cli/tui_pipeline.py index cd7fafd1..eec4e9e6 100644 --- a/composer/cli/tui_pipeline.py +++ b/composer/cli/tui_pipeline.py @@ -15,14 +15,13 @@ from typing import cast, Protocol -from composer.core.user import user_data_ns from composer.input.types import ModelOptions, RAGDBOptions, DEFAULT_RECURSION_LIMIT from composer.input.parsing import add_protocol_args -from composer.io.thread_logging import DEFAULT_META_NS, thread_logger, default_logging_ns +from composer.io.thread_logging import thread_logger, default_logging_ns from composer.spec.agent_index import agent_index_config_from_env from composer.rag.db import PostgreSQLRAGDatabase from composer.rag.models import get_model -from composer.workflow.services import llm_factory, standard_connections +from composer.workflow.services import standard_connections from composer.spec.service_host import ModelProvider from composer.kb.knowledge_base import DefaultEmbedder, DEFAULT_KB_NS from composer.spec.services import build_rag_tool_env @@ -128,12 +127,10 @@ async def _main() -> int: # Set up services. Natspec does not support model-swapping, so the heavy # and lite tiers collapse onto the single configured model. - model_factory = llm_factory(args) model = get_model() llm_provider = get_provider_for(options=args) - logging_ns = user_data_ns() + DEFAULT_META_NS run_id = uuid.uuid4().hex async with ( diff --git a/composer/input/files.py b/composer/input/files.py index 87aa5179..c2d8b186 100644 --- a/composer/input/files.py +++ b/composer/input/files.py @@ -23,16 +23,16 @@ import asyncio import hashlib -import io import mimetypes -import os import pathlib import zlib from abc import ABC, abstractmethod from dataclasses import dataclass from typing import Protocol, assert_never, Any, overload -from composer.llm.provider import ProviderKind +class ContentRenderer(Protocol): + def text_block(self, text: str, *, with_cache: bool) -> dict: ... + def file_block(self, file_id: str, *, with_cache: bool) -> dict: ... # --------------------------------------------------------------------------- # Protocols (the public surface) @@ -123,20 +123,14 @@ class InMemoryTextFile: basename: str string_contents: str - provider: ProviderKind + renderer: ContentRenderer @property def bytes_contents(self) -> bytes: return self.string_contents.encode("utf-8") def to_dict(self, with_cache: bool = False) -> dict: - to_ret : dict[str, Any] = {"type": "text", "text": self.string_contents} - if with_cache: - to_ret["cache_control"] = { - "type": "ephemeral", - "ttl": "5m" - } - return to_ret + return self.renderer.text_block(self.string_contents, with_cache=with_cache) def to_digest(self) -> str: return _bytes_digest(self.bytes_contents) @@ -156,33 +150,10 @@ class UploadedFile: basename: str contents: bytes digest: str - provider: ProviderKind + renderer: ContentRenderer def to_dict(self, with_cache: bool = False) -> dict: - match self.provider: - case "anthropic": - to_ret : dict[str, Any] = { - "type": "document", - "source": { - "type": "file", - "file_id": self.file_id, - }, - } - if with_cache: - to_ret["cache_control"] = { - "type": "ephemeral", - "ttl": "5m" - } - return to_ret - case "openai": - return { - "type": "file", - "file": { - "file_id": self.file_id, - }, - } - case _: - assert_never(self.provider) + return self.renderer.file_block(file_id=self.file_id, with_cache=with_cache) def to_digest(self) -> str: return self.digest @@ -278,8 +249,6 @@ def _file_data_impl( class FileUploader(Protocol): """Upload+dedup contract. Obtain via ``ModelProvider.uploader()`` (``composer.llm``).""" - provider: ProviderKind - async def upload_file_if_needed( self, file_path: str | pathlib.Path ) -> UploadedFile: ... @@ -301,7 +270,7 @@ async def document_from(self, src: Uploadable) -> Document: -class _UploaderBase(ABC): +class UploaderBase(ABC): """Shared upload-or-reuse logic. Subclasses supply :meth:`_upload_bytes` (the provider-specific API call) and set ``provider`` at construction so it's stamped onto returned @@ -312,7 +281,7 @@ class _UploaderBase(ABC): so we don't reupload a file whose bytes the account has already seen.""" - provider: ProviderKind + renderer: ContentRenderer @abstractmethod async def _upload_bytes( @@ -334,7 +303,7 @@ async def upload_file_if_needed( basename=data.basename, contents=data.raw_data, digest=data.digest, - provider=self.provider, + renderer=self.renderer, ) async def upload_text_file_if_needed( @@ -352,7 +321,7 @@ async def upload_text_file_if_needed( basename=data.basename, contents=data.raw_data, digest=data.digest, - provider=self.provider, + renderer=self.renderer, ) async def get_document( @@ -379,12 +348,12 @@ async def get_document( basename=data.basename, contents=data.raw_data, digest=data.digest, - provider=self.provider + renderer=self.renderer ) return InMemoryTextFile( basename=p.name, string_contents=data.raw_data.decode("utf-8"), - provider=self.provider, + renderer=self.renderer, ) async def upload_bytes_if_needed( @@ -400,13 +369,13 @@ async def upload_bytes_if_needed( basename=data.basename, contents=data.raw_data, digest=data.digest, - provider=self.provider + renderer=self.renderer ) def text_document_from(self, src: TextUploadable) -> TextDocument: """Rehydrate a text ``Uploadable`` into an inline ``TextDocument`` (no upload — text stays in-prompt for transcript debuggability).""" - return InMemoryTextFile(basename=src.basename, string_contents=src.string_contents, provider=self.provider) + return InMemoryTextFile(basename=src.basename, string_contents=src.string_contents, renderer=self.renderer) async def document_from(self, src: Uploadable) -> Document: """Rehydrate an ``Uploadable`` (e.g. an audit-restored handle) into a @@ -414,6 +383,6 @@ async def document_from(self, src: Uploadable) -> Document: binary goes through the Files API for the active provider.""" text = src.string_contents if text is not None: - return InMemoryTextFile(basename=src.basename, string_contents=text, provider=self.provider) + return InMemoryTextFile(basename=src.basename, string_contents=text, renderer=self.renderer) return await self.upload_bytes_if_needed(src.basename, src.bytes_contents) diff --git a/composer/llm/anthropic.py b/composer/llm/anthropic.py index 61b86bc6..4cf4cab4 100644 --- a/composer/llm/anthropic.py +++ b/composer/llm/anthropic.py @@ -5,14 +5,16 @@ from io import BytesIO from dataclasses import dataclass, field import asyncio +from functools import cache import anthropic -from composer.input.files import _UploaderBase +from composer.input.files import UploaderBase, ContentRenderer from composer.input.types import ModelConfiguration from composer.llm.provider import ( - ProviderKind, CacheLevel, _ListIter, NoSuchElementError, + CacheLevel, ProviderServiceBase, ProviderSpec ) +from .parsing import ListIter, NoSuchElementError if TYPE_CHECKING: from langchain_core.language_models.chat_models import BaseChatModel @@ -38,7 +40,7 @@ def _validate_model(s: str) -> TypeGuard[ClaudeModelNames]: _interleaved_pivot_version = (4, 5) def _model_parser(model_name: str) -> ModelFeatures: - stream = _ListIter(model_name.split("-")) + stream = ListIter(model_name.split("-")) parsing: Literal["claude", "model", "version"] = "claude" try: claude = stream.next() @@ -77,14 +79,43 @@ def matches(model: str) -> bool: # --- Files API uploader ---------------------------------------------------- +class AnthropicRenderer: + def text_block(self, text, *, with_cache: bool) -> dict: + to_ret : dict[str, Any] = {"type": "text", "text": text} + if with_cache: + to_ret["cache_control"] = { + "type": "ephemeral", + "ttl": "5m" + } + return to_ret + + def file_block(self, file_id, *, with_cache: bool) -> dict: + to_ret : dict[str, Any] = { + "type": "document", + "source": { + "type": "file", + "file_id": file_id, + }, + } + if with_cache: + to_ret["cache_control"] = { + "type": "ephemeral", + "ttl": "5m" + } + return to_ret + +@cache +def _get_service(): + return AnthropicService() + @dataclass -class AnthropicFileUploader(_UploaderBase): +class AnthropicFileUploader(UploaderBase): """``FileUploader`` impl backed by Anthropic's beta Files API.""" client: anthropic.AsyncAnthropic uploaded: dict[str, str] | None = None _seed_lock: asyncio.Lock = field(default_factory=asyncio.Lock) - provider: ProviderKind = "anthropic" + renderer: ContentRenderer = field(default_factory=AnthropicRenderer) async def _ensure_seeded(self) -> dict[str, str]: """Seed the dedup cache from the account's existing Files-API uploads on @@ -117,6 +148,16 @@ def lazy() -> "AnthropicFileUploader": # --- ModelProvider --------------------------------------------------------- +class AnthropicService(ProviderServiceBase): + def __init__(self): + from graphcore.tools.memory import anthropic_async_memory_tool + super().__init__( + anthropic_async_memory_tool, + AnthropicFileUploader.lazy + ) + + + @dataclass class AnthropicModelProvider: """``ModelProvider`` for Anthropic. Probes ``model_name`` once at @@ -126,11 +167,16 @@ class AnthropicModelProvider: model_name: str options: ModelConfiguration features: ModelFeatures - provider: ProviderKind = "anthropic" + + provider: AnthropicService = field(default_factory=_get_service) @staticmethod def create(model_name: str, options: ModelConfiguration) -> "AnthropicModelProvider": - return AnthropicModelProvider(model_name, options, _model_parser(model_name)) + return AnthropicModelProvider( + model_name, + options, + _model_parser(model_name), + ) def builder_for( self, *, cache_level: CacheLevel | None = None, disable_thinking: bool = False @@ -176,3 +222,8 @@ def builder_for( model_kwargs=model_kwargs, callbacks=[UsageCallback()], ) + +ANTHROPIC_SPEC = ProviderSpec( + matches=matches, + build=AnthropicModelProvider.create +) diff --git a/composer/llm/api.py b/composer/llm/api.py new file mode 100644 index 00000000..0656c2ab --- /dev/null +++ b/composer/llm/api.py @@ -0,0 +1,7 @@ +from typing import Callable +from dataclasses import dataclass +from .provider import ModelProvider + +from composer.input.types import ModelConfiguration + + diff --git a/composer/llm/openai.py b/composer/llm/openai.py index ccca867c..2bca9e34 100644 --- a/composer/llm/openai.py +++ b/composer/llm/openai.py @@ -14,18 +14,22 @@ from typing import Literal, TypeGuard, Any, TYPE_CHECKING import io from dataclasses import dataclass, field +from functools import cache import asyncio import openai -from composer.input.files import _UploaderBase +from composer.input.files import UploaderBase, FileUploader, ContentRenderer from composer.input.types import ModelConfiguration from composer.llm.provider import ( - ProviderKind, CacheLevel, _ListIter, NoSuchElementError, + CacheLevel, ProviderServiceBase, ProviderSpec ) +from .parsing import ListIter, NoSuchElementError if TYPE_CHECKING: from langchain_core.language_models.chat_models import BaseChatModel + from langchain_core.tools import BaseTool + from graphcore.tools.memory import AsyncPostgresBackend # --- model probing --------------------------------------------------------- @@ -67,7 +71,7 @@ def _parse_gpt_version(token: str) -> tuple[int, int]: return (int(token), 0) def _model_parser(model_name: str) -> OpenAIModelFeatures: - stream = _ListIter(model_name.split("-")) + stream = ListIter(model_name.split("-")) parsing: Literal["family", "version", "tier"] = "family" try: head = stream.next() @@ -121,18 +125,38 @@ def _reasoning_effort(thinking_tokens: int) -> Literal["low", "medium", "high"]: return "medium" return "high" +class OpenAIService(ProviderServiceBase): + def __init__(self): + from graphcore.tools.memory import openai_async_memory_tool + super().__init__( + openai_async_memory_tool, + OpenAIFileUploader.lazy + ) + +@dataclass +class OpenAIRenderer: + def text_block(self, text, *, with_cache: bool) -> dict: + to_ret : dict[str, Any] = {"type": "text", "text": text} + return to_ret + def file_block(self, file_id: str, *, with_cache: bool) -> dict: + return { + "type": "file", + "file": { + "file_id": file_id, + }, + } # --- Files API uploader ---------------------------------------------------- @dataclass -class OpenAIFileUploader(_UploaderBase): +class OpenAIFileUploader(UploaderBase): """``FileUploader`` impl backed by OpenAI's Files API (``purpose="user_data"``).""" client: openai.AsyncOpenAI uploaded: dict[str, str] | None = None _seed_lock: asyncio.Lock = field(default_factory=asyncio.Lock) - provider: ProviderKind = "openai" + provider: ContentRenderer = field(default_factory=OpenAIRenderer) async def _ensure_seeded(self) -> dict[str, str]: """Seed the dedup cache from the account's existing user-data uploads on @@ -162,6 +186,10 @@ async def _upload_bytes( def lazy() -> "OpenAIFileUploader": """A lazily-seeding uploader — no account file-list until first upload.""" return OpenAIFileUploader(client=openai.AsyncOpenAI()) + +@cache +def _openai_service(): + return OpenAIService() # --- ModelProvider --------------------------------------------------------- @@ -175,7 +203,7 @@ class OpenAIModelProvider: model_name: str options: ModelConfiguration features: OpenAIModelFeatures - provider: ProviderKind = "openai" + provider: OpenAIService = field(default_factory=_openai_service) @staticmethod def create(model_name: str, options: ModelConfiguration) -> "OpenAIModelProvider": @@ -207,3 +235,8 @@ def builder_for( max_retries=2, **kwargs, ) + +OPEN_AI_SPEC = ProviderSpec( + matches=matches, + build=OpenAIModelProvider.create +) diff --git a/composer/llm/parsing.py b/composer/llm/parsing.py new file mode 100644 index 00000000..b6458fb2 --- /dev/null +++ b/composer/llm/parsing.py @@ -0,0 +1,25 @@ +from dataclasses import dataclass, field + +class NoSuchElementError(RuntimeError): + pass + + +@dataclass +class ListIter[T]: + l: list[T] + ind: int = field(default=0) + + def has_next(self) -> bool: + return self.ind < len(self.l) + + def peek(self) -> T: + if not self.has_next(): + raise NoSuchElementError("Invalid state, no more elements") + return self.l[self.ind] + + def next(self) -> T: + if not self.has_next(): + raise NoSuchElementError("Invalid state, no more elements") + to_ret = self.l[self.ind] + self.ind += 1 + return to_ret diff --git a/composer/llm/provider.py b/composer/llm/provider.py index 8967cf75..84b58d22 100644 --- a/composer/llm/provider.py +++ b/composer/llm/provider.py @@ -9,16 +9,47 @@ the per-provider modules can import it without an import cycle. """ -from typing import Literal, Protocol, TYPE_CHECKING -from dataclasses import dataclass, field +from typing import Protocol, TYPE_CHECKING, Callable +from dataclasses import dataclass +from functools import cached_property import enum +from composer.input.files import FileUploader +from composer.input.types import ModelConfiguration +from abc import ABC if TYPE_CHECKING: from langchain_core.language_models.chat_models import BaseChatModel - - -type ProviderKind = Literal["anthropic", "openai"] - + from graphcore.tools.memory import AsyncPostgresBackend + from langchain_core.tools import BaseTool + +class ProviderService(Protocol): + def select_memory_tool( + self, backend: "AsyncPostgresBackend" + ) -> "BaseTool": + ... + + def uploader(self) -> FileUploader: + ... + +class ProviderServiceBase(ABC): + def __init__(self, + mem_fact: Callable[["AsyncPostgresBackend"], "BaseTool"], + uploader_fact: Callable[[], FileUploader] + ): + self.mem_fact = mem_fact + self.uploader_fact = uploader_fact + + @cached_property + def _uploader_prop(self) -> FileUploader: + return self.uploader_fact() + + def uploader(self) -> FileUploader: + return self._uploader_prop + + def select_memory_tool( + self, backend: "AsyncPostgresBackend" + ) -> "BaseTool": + return self.mem_fact(backend) class CacheLevel(enum.StrEnum): NONE = "none" @@ -33,33 +64,15 @@ class ModelProvider(Protocol): mints a chat model with the cache/thinking choice deferred to the call site.""" @property - def provider(self) -> ProviderKind: ... + def provider(self) -> ProviderService: ... def builder_for( self, *, cache_level: CacheLevel | None = None, disable_thinking: bool = False ) -> "BaseChatModel": ... - -class NoSuchElementError(RuntimeError): - pass - - -@dataclass -class _ListIter[T]: - l: list[T] - ind: int = field(default=0) - - def has_next(self) -> bool: - return self.ind < len(self.l) - - def peek(self) -> T: - if not self.has_next(): - raise NoSuchElementError("Invalid state, no more elements") - return self.l[self.ind] - - def next(self) -> T: - if not self.has_next(): - raise NoSuchElementError("Invalid state, no more elements") - to_ret = self.l[self.ind] - self.ind += 1 - return to_ret +@dataclass(frozen=True) +class ProviderSpec: + """One row of the provider registry: a name predicate, the provider kind it + maps to, and the factory that builds the provider's ``ModelProvider``.""" + matches: Callable[[str], bool] + build: Callable[[str, ModelConfiguration], ModelProvider] diff --git a/composer/llm/registry.py b/composer/llm/registry.py index 871122fe..51212998 100644 --- a/composer/llm/registry.py +++ b/composer/llm/registry.py @@ -8,62 +8,45 @@ """ from dataclasses import dataclass -from typing import Callable, Protocol, TYPE_CHECKING, overload, cast +from typing import overload, cast +from functools import cache +import importlib.metadata from composer.input.types import ModelConfiguration, ModelOptionsBase, TieredModelOptions -from composer.input.files import FileUploader -from composer.llm.provider import ProviderKind, CacheLevel, ModelProvider -from composer.llm import anthropic as _anthropic -from composer.llm import openai as _openai - -if TYPE_CHECKING: - from langchain_core.language_models.chat_models import BaseChatModel - - -@dataclass(frozen=True) -class ProviderSpec: - """One row of the provider registry: a name predicate, the provider kind it - maps to, and the factory that builds the provider's ``ModelProvider``.""" - matches: Callable[[str], bool] - kind: ProviderKind - build: Callable[[str, ModelConfiguration], ModelProvider] - - -_PROVIDERS: list[ProviderSpec] = [ - ProviderSpec( - matches=_anthropic.matches, - kind="anthropic", - build=_anthropic.AnthropicModelProvider.create, - ), - ProviderSpec( - matches=_openai.matches, - kind="openai", - build=_openai.OpenAIModelProvider.create, - ), -] - +from composer.llm.provider import ProviderService, ModelProvider +from .provider import ProviderSpec + +LLM_PROVIDER_GROUP = "certora.autoprove.llm_provider" + +@cache +def _loader_providers() -> list[ProviderSpec]: + to_ret : list[ProviderSpec] = [] + for ep in importlib.metadata.entry_points( + group=LLM_PROVIDER_GROUP + ): + prov = ep.load() + if not isinstance(prov, ProviderSpec): + raise ValueError(f"Could not load provider backend: {ep.name} with {ep.module}.{ep.value}") + to_ret.append(prov) + return to_ret def _lookup(model: str) -> ProviderSpec: lowered = model.lower() - for spec in _PROVIDERS: + for spec in _loader_providers(): if spec.matches(lowered): return spec raise ValueError( f"Unrecognized model {model!r}: cannot determine its provider. Add a " - f"ProviderSpec to composer.llm.registry._PROVIDERS when introducing a " + f"ProviderSpec to the `{LLM_PROVIDER_GROUP}` importlib entrypoint when introducing a " f"new model family." ) -def provider_for(model: str) -> ProviderKind: - """Map a model identifier to its provider family via the registry.""" - return _lookup(model).kind - @dataclass(kw_only=True) class TieredProviders: lite: ModelProvider heavy: ModelProvider - provider_kind: ProviderKind + provider_service: ProviderService @overload def get_provider_for(*, model_name: str, options: ModelConfiguration) -> ModelProvider: @@ -94,41 +77,6 @@ def get_provider_for( assert tiered is not None lite_model = _lookup(tiered.lite_model).build(tiered.lite_model, tiered) heavy_model = _lookup(tiered.heavy_model).build(tiered.heavy_model, tiered) - if lite_model.provider != heavy_model.provider: + if type(lite_model.provider) is not type(heavy_model.provider): raise ValueError(f"Cannot use different model providers for heavy and lite models: {tiered.lite_model} vs {tiered.heavy_model}") - return TieredProviders(lite=lite_model, heavy=heavy_model, provider_kind=lite_model.provider) - -def uploader_for(provider: ProviderKind) -> FileUploader: - """Construct the lazily-seeding Files-API uploader for ``provider``.""" - match provider: - case "anthropic": - return _anthropic.AnthropicFileUploader.lazy() - case "openai": - return _openai.OpenAIFileUploader.lazy() - - -class LLMFactory(Protocol): - def __call__( - self, - model_name: str, - *, - cache_level: CacheLevel | None = None, - disable_thinking: bool = False, - ) -> "BaseChatModel": ... - - -def llm_factory(options: ModelConfiguration) -> LLMFactory: - """A model-name → chat-model factory bound to ``options``. The tiering layer - (``ModelProvider`` in ``spec/service_host.py``) calls this per tier with the - heavy/lite model name; each call resolves the provider and defers the - cache/thinking choice to ``builder_for``.""" - def build( - model_name: str, - *, - cache_level: CacheLevel | None = None, - disable_thinking: bool = False, - ) -> "BaseChatModel": - return get_provider_for(model_name=model_name, options=options).builder_for( - cache_level=cache_level, disable_thinking=disable_thinking - ) - return build + return TieredProviders(lite=lite_model, heavy=heavy_model, provider_service=lite_model.provider) diff --git a/composer/natreq/extractor.py b/composer/natreq/extractor.py index 0be34a5f..9e616e8b 100644 --- a/composer/natreq/extractor.py +++ b/composer/natreq/extractor.py @@ -22,11 +22,9 @@ from composer.rag.db import ComposerRAGDB, rag_context from composer.rag.models import get_model from composer.workflow.services import checkpointer_context -from composer.workflow.provider import ProviderKind from composer.tools.search import cvl_manual_search from composer.tools.thinking import RoughDraftState, get_rough_draft_tools from composer.templates.loader import load_jinja_template -from composer.human.types import HumanInteractionType from composer.io.protocol import IOHandler from composer.io.context import with_handler, run_graph from composer.io.event_handler import NullEventHandler diff --git a/composer/pipeline/cli.py b/composer/pipeline/cli.py index ccdd16c5..46e0b714 100644 --- a/composer/pipeline/cli.py +++ b/composer/pipeline/cli.py @@ -172,7 +172,7 @@ async def cli_pipeline[P: enum.Enum, H]( ) async with ( - standard_connections(provider=tiered.provider_kind, embedder=DefaultEmbedder(model)) as conns, + standard_connections(provider=tiered.provider_service, embedder=DefaultEmbedder(model)) as conns, async_tool_context(), thread_logger(conns.store, { "root_thread_id": thread_id, diff --git a/composer/spec/services.py b/composer/spec/services.py index 35769fa2..c7a00ba1 100644 --- a/composer/spec/services.py +++ b/composer/spec/services.py @@ -8,7 +8,6 @@ from composer.spec.service_host import ModelProvider, PureServiceHost, Sort from composer.spec.cvl_research import indexed_cvl_research_tool, CVL_RESEARCH_BASE_DOC from composer.tools.search import cvl_manual_tools -from composer.workflow.provider import ProviderKind from composer.kb.knowledge_base import kb_tools from composer.spec.agent_index import AgentIndex, AgentIndexConfig, RetrieveDocumentTool diff --git a/composer/spec/source/author.py b/composer/spec/source/author.py index d0bf3cd3..58aa47bb 100644 --- a/composer/spec/source/author.py +++ b/composer/spec/source/author.py @@ -26,7 +26,7 @@ from pathlib import Path from composer.spec.gen_types import CVLResource, TypedTemplate, import_statement_for from composer.spec.service_host import ServiceHost -from composer.workflow.services import CacheLevel +from composer.llm.provider import CacheLevel from langgraph.types import Command diff --git a/composer/testing/harness_tape.py b/composer/testing/harness_tape.py index ef92bf73..c7eb33e8 100644 --- a/composer/testing/harness_tape.py +++ b/composer/testing/harness_tape.py @@ -6,7 +6,6 @@ from langchain_core.language_models.fake_chat_models import ( FakeMessagesListChatModel, ) -from composer.llm.provider import ProviderKind from langchain_core.prompt_values import PromptValue from langchain_core.tools import BaseTool from langchain_core.messages import BaseMessage, AIMessage @@ -119,32 +118,50 @@ async def ainvoke( class _DummyUploader: - """A ``FileUploader`` stand-in that never touches a Files API: every input is - read from disk and returned as an in-memory text document. Installed under the - harness so a taped run does no real uploads — the codegen path uploads the spec - + interface via ``upload_text_file_if_needed``, which would otherwise hit the - live Files API.""" + """A ``FileUploader`` stand-in that never touches a Files API: every input — + a path (``upload_*``/``get_document``) or an ``Uploadable`` handle + (``document_from``/``text_document_from``) — is returned as an in-memory text + document. Installed under the harness so a taped run does no real uploads, + which would otherwise hit the live Files API.""" async def upload_text_file_if_needed(self, file_path: Any) -> Any: - return self._inline(file_path) + return self._inline_path(file_path) async def upload_file_if_needed(self, file_path: Any) -> Any: - return self._inline(file_path) + return self._inline_path(file_path) async def get_document(self, path: Any) -> Any: import os - return self._inline(path) if os.path.isfile(str(path)) else None + return self._inline_path(path) if os.path.isfile(str(path)) else None - @staticmethod - def _inline(path: Any) -> Any: + def text_document_from(self, src: Any) -> Any: + return self._inline_src(src) + + async def document_from(self, src: Any) -> Any: + return self._inline_src(src) + + @classmethod + def _inline_path(cls, path: Any) -> Any: import os from pathlib import Path - from composer.input.files import InMemoryTextFile p = str(path) + return cls._doc(os.path.basename(p), Path(p).read_text(encoding="utf-8")) + + @classmethod + def _inline_src(cls, src: Any) -> Any: + text = src.string_contents + if text is None: + text = src.bytes_contents.decode("utf-8", errors="replace") + return cls._doc(src.basename, text) + + @staticmethod + def _doc(basename: str, text: str) -> Any: + from composer.input.files import InMemoryTextFile + from composer.llm.anthropic import AnthropicRenderer return InMemoryTextFile( - basename=os.path.basename(p), - string_contents=Path(p).read_text(encoding="utf-8"), - provider="anthropic" + basename=basename, + string_contents=text, + renderer=AnthropicRenderer(), ) @@ -189,8 +206,21 @@ def install_fake_llm(fake: Any) -> None: import composer.llm.registry as registry import composer.workflow.services as services + from composer.llm.provider import ProviderServiceBase + + class _FakeService(ProviderServiceBase): + """The real anthropic memory tool (backed by the harness's Postgres + memory backend), paired with the dummy uploader so no taped run touches + a live Files API. ``ProviderServiceBase`` memoizes the uploader, so a + single ``_DummyUploader`` is shared across the CLI pre-upload and the + in-workflow uploads.""" + + def __init__(self) -> None: + from graphcore.tools.memory import anthropic_async_memory_tool + super().__init__(anthropic_async_memory_tool, _DummyUploader) + class _FakeProvider: - provider : ProviderKind = "anthropic" + provider = _FakeService() def builder_for(self, *, cache_level: Any = None, disable_thinking: bool = False) -> Any: return _current_fake() @@ -201,11 +231,10 @@ def _fake_get_provider_for( *, model_name: Any = None, options: Any = None, tiered: Any = None ) -> Any: if tiered is not None: - return registry.TieredProviders(lite=fp, heavy=fp, provider_kind="anthropic") + return registry.TieredProviders(lite=fp, heavy=fp, provider_service=fp.provider) return fp registry.get_provider_for = _fake_get_provider_for - registry.uploader_for = lambda _provider: _DummyUploader() services.create_llm = lambda args: _current_fake() services.create_llm_base = lambda args: _current_fake() _llm_seams_patched = True diff --git a/composer/workflow/provider.py b/composer/workflow/provider.py deleted file mode 100644 index 95d73588..00000000 --- a/composer/workflow/provider.py +++ /dev/null @@ -1,11 +0,0 @@ -"""Back-compat shim. - -Provider classification + the model feature matrices moved to ``composer.llm``. -New code should import ``ProviderKind`` from ``composer.llm.provider`` and -``provider_for`` / ``get_provider_for`` from ``composer.llm.registry`` directly. -""" - -from composer.llm.provider import ProviderKind -from composer.llm.registry import provider_for - -__all__ = ["ProviderKind", "provider_for"] diff --git a/composer/workflow/services.py b/composer/workflow/services.py index 76bcfdc5..b7fb5a0f 100644 --- a/composer/workflow/services.py +++ b/composer/workflow/services.py @@ -1,5 +1,5 @@ import psycopg -from typing import Any, Callable, TypedDict, Literal, overload, AsyncContextManager, TYPE_CHECKING, assert_never +from typing import Any, Callable, TypedDict, Literal, overload, AsyncContextManager, TYPE_CHECKING from typing_extensions import TypeVar import inspect import os @@ -40,105 +40,12 @@ from composer.input.types import ModelOptions, ModelOptionsBase from composer.input.files import FileUploader -from composer.llm.provider import ProviderKind, CacheLevel -from composer.llm.registry import get_provider_for, uploader_for, llm_factory, LLMFactory +from composer.llm.provider import ProviderService +from composer.llm.registry import get_provider_for T = TypeVar("T") -def _adapt_async(obj: T, pairs: list[tuple[str, str]]) -> T: - """ - Patch async methods to forward to their sync counterparts. - - Args: - obj: Object to patch - pairs: List of (async_name, sync_name) tuples - - Raises: - AttributeError: If method names don't exist on obj - TypeError: If async method is not a coroutine or sync method is a coroutine - ValueError: If method signatures don't match - """ - for async_name, sync_name in pairs: - # Step 1: Fetch attributes - try: - async_method = getattr(obj, async_name) - except AttributeError: - raise AttributeError( - f"Object {obj} does not have async method '{async_name}'" - ) - - try: - sync_method = getattr(obj, sync_name) - except AttributeError: - raise AttributeError( - f"Object {obj} does not have sync method '{sync_name}'" - ) - - # Step 2: Verify that async_method is a coroutine function - if not inspect.iscoroutinefunction(async_method) and not inspect.isasyncgenfunction(async_method): - raise TypeError( - f"Method '{async_name}' is not a coroutine function" - ) - - # Verify that sync_method is NOT a coroutine function - if inspect.iscoroutinefunction(sync_method): - raise TypeError( - f"Method '{sync_name}' is a coroutine function but should be sync" - ) - - # Get signatures - async_sig = inspect.signature(async_method) - sync_sig = inspect.signature(sync_method) - - # Compare parameters (names and annotations) - async_params = list(async_sig.parameters.values()) - sync_params = list(sync_sig.parameters.values()) - - if len(async_params) != len(sync_params): - raise ValueError( - f"Parameter count mismatch: {async_name} has {len(async_params)} " - f"parameters, {sync_name} has {len(sync_params)}" - ) - - for async_param, sync_param in zip(async_params, sync_params): - if async_param.name != sync_param.name: - raise ValueError( - f"Parameter name mismatch: {async_name} has '{async_param.name}', " - f"{sync_name} has '{sync_param.name}'" - ) - - if async_param.annotation != sync_param.annotation: - raise ValueError( - f"Parameter annotation mismatch for '{async_param.name}': " - f"{async_name} has {async_param.annotation}, " - f"{sync_name} has {sync_param.annotation}" - ) - - if async_param.default != sync_param.default: - raise ValueError( - f"Parameter default mismatch for '{async_param.name}': " - f"{async_name} has {async_param.default}, " - f"{sync_name} has {sync_param.default}" - ) - - # Step 3: Create wrapper that forwards to sync implementation - def make_wrapper(sync_fn: Callable) -> Callable: - async def async_wrapper(*args, **kwargs): - # Call the sync function - return sync_fn(*args, **kwargs) - - # Preserve the original signature - setattr(async_wrapper, "__signature__", inspect.signature(sync_fn)) - async_wrapper.__name__ = sync_fn.__name__ - async_wrapper.__doc__ = sync_fn.__doc__ - - return async_wrapper - - # Patch the object - new_async_method = make_wrapper(sync_method) - setattr(obj, async_name, new_async_method) - return obj # Bound each connect attempt; give getconn a long window to keep retrying so a # slow first connection doesn't fail the run. @@ -314,26 +221,6 @@ async def store_context() -> AsyncIterator[AsyncPostgresStore]: from typing_extensions import deprecated -@deprecated("Use async code") -def get_checkpointer() -> PostgresSaver: - conn = _get_composer_connection( - **_DATABASE_CONFIGS["checkpoint"], - autocommit=True, - row_factory=dict_row - ) - checkpointer = _adapt_async( - PostgresSaver(conn), - [("aget", "get"), - ("aput", "put"), - ("aget_tuple", "get_tuple"), - ("alist", "list"), - ("adelete_thread", "delete_thread"), - ("aput_writes", "put_writes") - ] - ) - checkpointer.setup() - return checkpointer - async def get_async_checkpointer() -> AsyncPostgresSaver: conn = await _get_async_composer_pool( **_DATABASE_CONFIGS["checkpoint"], @@ -344,17 +231,6 @@ async def get_async_checkpointer() -> AsyncPostgresSaver: await checkpointer.setup() return checkpointer -@deprecated("Use async code") -def get_store() -> PostgresStore: - conn = _get_composer_connection( - **_DATABASE_CONFIGS["store"], - autocommit=True, - row_factory=dict_row - ) - store = PostgresStore(conn) - store.setup() - return store - async def get_async_store() -> AsyncPostgresStore: conn = await _get_async_composer_pool( **_DATABASE_CONFIGS["store"], @@ -467,7 +343,6 @@ class StandardConnections: store: AsyncPostgresStore memory: "Callable[[str], BaseTool]" uploader: FileUploader - provider: ProviderKind @dataclass class IndexedConnections(StandardConnections): @@ -475,7 +350,7 @@ class IndexedConnections(StandardConnections): def _memory_tool_factory( - provider: ProviderKind, + provider: ProviderService, backend_factory: Callable[[str], AsyncPostgresBackend], ) -> Callable[[str], BaseTool]: """Wrap a backend factory into a provider-aware tool factory. @@ -483,36 +358,28 @@ def _memory_tool_factory( Closes over the provider so the rest of the codebase only sees ``Callable[[ns], BaseTool]`` — the memory tool's shape is the factory's secret.""" - from graphcore.tools.memory import anthropic_async_memory_tool, openai_async_memory_tool - - match provider: - case "anthropic": - return lambda ns: anthropic_async_memory_tool(backend_factory(ns)) - case "openai": - return lambda ns: openai_async_memory_tool(backend_factory(ns)) - case _: - assert_never(provider) + return lambda ns: provider.select_memory_tool(backend_factory(ns)) @overload def standard_connections( - *, provider: ProviderKind, + *, provider: ProviderService, ) -> AsyncContextManager[StandardConnections]: ... @overload def standard_connections( - *, provider: ProviderKind, embedder: Embeddings + *, provider: ProviderService, embedder: Embeddings ) -> AsyncContextManager[IndexedConnections]: ... @asynccontextmanager async def standard_connections( *, - provider: ProviderKind, + provider: ProviderService, embedder: Embeddings | None = None, ) -> AsyncIterator[StandardConnections | IndexedConnections]: - uploader = uploader_for(provider) + uploader = provider.uploader() async with ( checkpointer_context() as check, memory_backend_context() as mem, @@ -527,7 +394,6 @@ async def standard_connections( store=store, memory=memory, uploader=uploader, - provider=provider, ) return yield StandardConnections( @@ -535,5 +401,4 @@ async def standard_connections( store=store, memory=memory, uploader=uploader, - provider=provider, ) diff --git a/pyproject.toml b/pyproject.toml index f999a88f..f2fbd53d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -166,3 +166,9 @@ certora_autosetup = [ "**/*.j2", "**/*.json", ] + + + +[project.entry-points."certora.autoprove.llm_provider"] +openai = "composer.llm.openai:OPEN_AI_SPEC" +anthropic = "composer.llm.anthropic:ANTHROPIC_SPEC" diff --git a/sanity_analyzer/analysis.py b/sanity_analyzer/analysis.py index 2aed973e..e5cccdcb 100644 --- a/sanity_analyzer/analysis.py +++ b/sanity_analyzer/analysis.py @@ -1,4 +1,4 @@ -from typing import NotRequired, assert_never +from typing import NotRequired import pathlib import uuid @@ -13,14 +13,13 @@ from composer.rag.models import get_model from composer.tools.search import cvl_manual_search from composer.templates.loader import load_jinja_template -from composer.workflow.services import create_llm, get_memory -from composer.workflow.provider import provider_for +from composer.workflow.services import create_llm, get_async_memory +from composer.llm.registry import get_provider_for from composer.tools.thinking import get_rough_draft_tools, RoughDraftState -from graphcore.tools.memory import anthropic_memory_tool, openai_memory_tool from graphcore.tools.vfs import fs_tools -from graphcore.graph import build_workflow, FlowInput, build_async_workflow +from graphcore.graph import build_workflow, FlowInput from graphcore.tools.results import result_tool_generator from sanity_analyzer.types import SanityAnalysisArgs @@ -249,18 +248,12 @@ async def async_analyze(args: SanityAnalysisArgs) -> SanityAnalysisResult | None if args.thread_id is None: print(f"Chose thread id: {tid}") - provider = provider_for(args.model) + provider = get_provider_for(model_name=args.model, options=args) rag_db = await get_rag_db(args.rag_db, model=get_model()) tools = [cvl_manual_search(rag_db), sanity_analysis_output_tool, *get_rough_draft_tools(SanityState), *v_tools] if args.memory_tool: - match provider: - case "anthropic": - mem_factory = anthropic_memory_tool - case "openai": - mem_factory = openai_memory_tool - case _: - assert_never(provider) - tools.append(mem_factory(get_memory(tid, "sanity"))) + mem_factory = provider.provider.select_memory_tool + tools.append(mem_factory(await get_async_memory(tid))) llm = create_llm(args) From 0f9624190f748b9b875c6e5c24ca141665212066 Mon Sep 17 00:00:00 2001 From: John Toman Date: Tue, 21 Jul 2026 12:43:20 -0700 Subject: [PATCH 2/4] fix tests --- tests/test_design_doc_finder.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/tests/test_design_doc_finder.py b/tests/test_design_doc_finder.py index 17e2aa84..f2ed53e0 100644 --- a/tests/test_design_doc_finder.py +++ b/tests/test_design_doc_finder.py @@ -16,7 +16,7 @@ integration tape. """ -from typing import Any, cast, Literal +from typing import Any, cast from dataclasses import dataclass @@ -35,9 +35,10 @@ from langgraph.checkpoint.memory import InMemorySaver +from composer.llm.anthropic import AnthropicRenderer, _get_service from composer.input.files import InMemoryTextFile from composer.spec.context import WorkflowContext, SourceFields -from composer.spec.service_host import ModelProvider, CoreModelProvider +from composer.spec.service_host import ModelProvider from composer.spec.util import FS_FORBIDDEN_READ from composer.templates.loader import load_jinja_template from composer.ui.autoprove_app import AutoProvePhase @@ -92,7 +93,7 @@ async def get_document(self, path: Any) -> InMemoryTextFile | None: p = pathlib.Path(path) if not p.is_file(): return None - return InMemoryTextFile(basename=p.name, string_contents=p.read_text(), provider="anthropic") + return InMemoryTextFile(basename=p.name, string_contents=p.read_text(), renderer=AnthropicRenderer()) def _source( @@ -321,7 +322,7 @@ class FakeModelFactory: @property def provider(self): - return "anthropic" + return _get_service() def builder_for(self, *args, **kwargs): return self.fake From ce725edc076a52a22898028839fbdc6fb40c9b21 Mon Sep 17 00:00:00 2001 From: John Toman Date: Mon, 27 Jul 2026 14:24:53 -0700 Subject: [PATCH 3/4] CR --- composer/llm/anthropic.py | 2 +- composer/llm/api.py | 7 ------- composer/llm/{parsing.py => list_iter.py} | 0 composer/llm/openai.py | 4 ++-- 4 files changed, 3 insertions(+), 10 deletions(-) delete mode 100644 composer/llm/api.py rename composer/llm/{parsing.py => list_iter.py} (100%) diff --git a/composer/llm/anthropic.py b/composer/llm/anthropic.py index 4cf4cab4..d3fb2d2e 100644 --- a/composer/llm/anthropic.py +++ b/composer/llm/anthropic.py @@ -14,7 +14,7 @@ from composer.llm.provider import ( CacheLevel, ProviderServiceBase, ProviderSpec ) -from .parsing import ListIter, NoSuchElementError +from .list_iter import ListIter, NoSuchElementError if TYPE_CHECKING: from langchain_core.language_models.chat_models import BaseChatModel diff --git a/composer/llm/api.py b/composer/llm/api.py deleted file mode 100644 index 0656c2ab..00000000 --- a/composer/llm/api.py +++ /dev/null @@ -1,7 +0,0 @@ -from typing import Callable -from dataclasses import dataclass -from .provider import ModelProvider - -from composer.input.types import ModelConfiguration - - diff --git a/composer/llm/parsing.py b/composer/llm/list_iter.py similarity index 100% rename from composer/llm/parsing.py rename to composer/llm/list_iter.py diff --git a/composer/llm/openai.py b/composer/llm/openai.py index 2bca9e34..699cf9d4 100644 --- a/composer/llm/openai.py +++ b/composer/llm/openai.py @@ -24,7 +24,7 @@ from composer.llm.provider import ( CacheLevel, ProviderServiceBase, ProviderSpec ) -from .parsing import ListIter, NoSuchElementError +from .list_iter import ListIter, NoSuchElementError if TYPE_CHECKING: from langchain_core.language_models.chat_models import BaseChatModel @@ -156,7 +156,7 @@ class OpenAIFileUploader(UploaderBase): client: openai.AsyncOpenAI uploaded: dict[str, str] | None = None _seed_lock: asyncio.Lock = field(default_factory=asyncio.Lock) - provider: ContentRenderer = field(default_factory=OpenAIRenderer) + renderer: ContentRenderer = field(default_factory=OpenAIRenderer) async def _ensure_seeded(self) -> dict[str, str]: """Seed the dedup cache from the account's existing user-data uploads on From b7b7c3efab1bcc3a7639e7965e2123c1718838a0 Mon Sep 17 00:00:00 2001 From: John Toman Date: Mon, 27 Jul 2026 15:35:48 -0700 Subject: [PATCH 4/4] Property caching ttl in document --- composer/input/files.py | 18 +++++++++------- composer/llm/anthropic.py | 37 +++++++++++++++++--------------- composer/llm/openai.py | 13 ++++++----- composer/llm/provider.py | 10 +++------ composer/llm/types.py | 6 ++++++ composer/spec/prop_inference.py | 3 ++- composer/testing/harness_tape.py | 4 ++-- composer/testing/record_tape.py | 3 ++- 8 files changed, 51 insertions(+), 43 deletions(-) create mode 100644 composer/llm/types.py diff --git a/composer/input/files.py b/composer/input/files.py index c2d8b186..49b88aae 100644 --- a/composer/input/files.py +++ b/composer/input/files.py @@ -28,11 +28,13 @@ import zlib from abc import ABC, abstractmethod from dataclasses import dataclass -from typing import Protocol, assert_never, Any, overload +from typing import Protocol, overload + +from composer.llm.types import CacheLevel class ContentRenderer(Protocol): - def text_block(self, text: str, *, with_cache: bool) -> dict: ... - def file_block(self, file_id: str, *, with_cache: bool) -> dict: ... + def text_block(self, text: str, *, cache_level: CacheLevel = CacheLevel.NONE) -> dict: ... + def file_block(self, file_id: str, *, cache_level: CacheLevel = CacheLevel.NONE) -> dict: ... # --------------------------------------------------------------------------- # Protocols (the public surface) @@ -65,7 +67,7 @@ def string_contents(self) -> str: ... class Document(Uploadable, Protocol): """A piece of content destined for an LLM message.""" - def to_dict(self, with_cache: bool = False) -> dict: ... + def to_dict(self, cache_level: CacheLevel = CacheLevel.NONE) -> dict: ... def to_digest(self) -> str: ... @@ -129,8 +131,8 @@ class InMemoryTextFile: def bytes_contents(self) -> bytes: return self.string_contents.encode("utf-8") - def to_dict(self, with_cache: bool = False) -> dict: - return self.renderer.text_block(self.string_contents, with_cache=with_cache) + def to_dict(self, cache_level: CacheLevel = CacheLevel.NONE) -> dict: + return self.renderer.text_block(self.string_contents, cache_level=cache_level) def to_digest(self) -> str: return _bytes_digest(self.bytes_contents) @@ -152,8 +154,8 @@ class UploadedFile: digest: str renderer: ContentRenderer - def to_dict(self, with_cache: bool = False) -> dict: - return self.renderer.file_block(file_id=self.file_id, with_cache=with_cache) + def to_dict(self, cache_level: CacheLevel = CacheLevel.NONE) -> dict: + return self.renderer.file_block(file_id=self.file_id, cache_level=cache_level) def to_digest(self) -> str: return self.digest diff --git a/composer/llm/anthropic.py b/composer/llm/anthropic.py index d3fb2d2e..33fe1a30 100644 --- a/composer/llm/anthropic.py +++ b/composer/llm/anthropic.py @@ -12,8 +12,9 @@ from composer.input.files import UploaderBase, ContentRenderer from composer.input.types import ModelConfiguration from composer.llm.provider import ( - CacheLevel, ProviderServiceBase, ProviderSpec + ProviderServiceBase, ProviderSpec ) +from .types import CacheLevel from .list_iter import ListIter, NoSuchElementError if TYPE_CHECKING: @@ -76,20 +77,28 @@ def _model_parser(model_name: str) -> ModelFeatures: def matches(model: str) -> bool: return model.split("-", 1)[0] == "claude" +def level_to_ttl(c: CacheLevel) -> str | None: + match c: + case CacheLevel.NONE: + return None + case CacheLevel.SHORT: + return "5m" + case CacheLevel.LONG: + return "1h" # --- Files API uploader ---------------------------------------------------- class AnthropicRenderer: - def text_block(self, text, *, with_cache: bool) -> dict: + def text_block(self, text: str, *, cache_level: CacheLevel = CacheLevel.NONE) -> dict: to_ret : dict[str, Any] = {"type": "text", "text": text} - if with_cache: + if (ttl := level_to_ttl(cache_level)) is not None: to_ret["cache_control"] = { "type": "ephemeral", - "ttl": "5m" + "ttl": ttl } return to_ret - - def file_block(self, file_id, *, with_cache: bool) -> dict: + + def file_block(self, file_id: str, *, cache_level: CacheLevel = CacheLevel.NONE) -> dict: to_ret : dict[str, Any] = { "type": "document", "source": { @@ -97,10 +106,10 @@ def file_block(self, file_id, *, with_cache: bool) -> dict: "file_id": file_id, }, } - if with_cache: + if (ttl := level_to_ttl(cache_level)) is not None: to_ret["cache_control"] = { "type": "ephemeral", - "ttl": "5m" + "ttl": ttl } return to_ret @@ -179,7 +188,7 @@ def create(model_name: str, options: ModelConfiguration) -> "AnthropicModelProvi ) def builder_for( - self, *, cache_level: CacheLevel | None = None, disable_thinking: bool = False + self, *, cache_level: CacheLevel = CacheLevel.NONE, disable_thinking: bool = False ) -> "BaseChatModel": from langchain_anthropic import ChatAnthropic from composer.diagnostics.usage_callback import UsageCallback @@ -199,15 +208,9 @@ def builder_for( if opts.memory_tool: betas.append("context-management-2025-06-27") - match cache_level: - case CacheLevel.SHORT: - ttl = "5m" - case CacheLevel.LONG: - ttl = "1h" - case None | CacheLevel.NONE: - ttl = None + ttl = level_to_ttl(cache_level) model_kwargs = ( - {"cache_control": {"type": "ephemeral", "ttl": ttl}} if ttl else {} + {"cache_control": {"type": "ephemeral", "ttl": ttl}} if ttl is not None else {} ) return ChatAnthropic( diff --git a/composer/llm/openai.py b/composer/llm/openai.py index 699cf9d4..46f2a1fb 100644 --- a/composer/llm/openai.py +++ b/composer/llm/openai.py @@ -19,17 +19,16 @@ import openai -from composer.input.files import UploaderBase, FileUploader, ContentRenderer +from composer.input.files import UploaderBase, ContentRenderer from composer.input.types import ModelConfiguration from composer.llm.provider import ( - CacheLevel, ProviderServiceBase, ProviderSpec + ProviderServiceBase, ProviderSpec ) +from .types import CacheLevel from .list_iter import ListIter, NoSuchElementError if TYPE_CHECKING: from langchain_core.language_models.chat_models import BaseChatModel - from langchain_core.tools import BaseTool - from graphcore.tools.memory import AsyncPostgresBackend # --- model probing --------------------------------------------------------- @@ -135,10 +134,10 @@ def __init__(self): @dataclass class OpenAIRenderer: - def text_block(self, text, *, with_cache: bool) -> dict: + def text_block(self, text: str, *, cache_level: CacheLevel = CacheLevel.NONE) -> dict: to_ret : dict[str, Any] = {"type": "text", "text": text} return to_ret - def file_block(self, file_id: str, *, with_cache: bool) -> dict: + def file_block(self, file_id: str, *, cache_level: CacheLevel = CacheLevel.NONE) -> dict: return { "type": "file", "file": { @@ -210,7 +209,7 @@ def create(model_name: str, options: ModelConfiguration) -> "OpenAIModelProvider return OpenAIModelProvider(model_name, options, _model_parser(model_name)) def builder_for( - self, *, cache_level: CacheLevel | None = None, disable_thinking: bool = False + self, *, cache_level: CacheLevel = CacheLevel.NONE, disable_thinking: bool = False ) -> "BaseChatModel": from langchain_openai import ChatOpenAI diff --git a/composer/llm/provider.py b/composer/llm/provider.py index 84b58d22..3814e04d 100644 --- a/composer/llm/provider.py +++ b/composer/llm/provider.py @@ -12,11 +12,12 @@ from typing import Protocol, TYPE_CHECKING, Callable from dataclasses import dataclass from functools import cached_property -import enum from composer.input.files import FileUploader from composer.input.types import ModelConfiguration +from .types import CacheLevel from abc import ABC + if TYPE_CHECKING: from langchain_core.language_models.chat_models import BaseChatModel from graphcore.tools.memory import AsyncPostgresBackend @@ -51,11 +52,6 @@ def select_memory_tool( ) -> "BaseTool": return self.mem_fact(backend) -class CacheLevel(enum.StrEnum): - NONE = "none" - SHORT = "short" - LONG = "long" - class ModelProvider(Protocol): """A provider-specific LLM backend, bound to one model. @@ -67,7 +63,7 @@ class ModelProvider(Protocol): def provider(self) -> ProviderService: ... def builder_for( - self, *, cache_level: CacheLevel | None = None, disable_thinking: bool = False + self, *, cache_level: CacheLevel = CacheLevel.NONE, disable_thinking: bool = False ) -> "BaseChatModel": ... @dataclass(frozen=True) diff --git a/composer/llm/types.py b/composer/llm/types.py new file mode 100644 index 00000000..ae9574ca --- /dev/null +++ b/composer/llm/types.py @@ -0,0 +1,6 @@ +import enum + +class CacheLevel(enum.StrEnum): + NONE = "none" + SHORT = "short" + LONG = "long" diff --git a/composer/spec/prop_inference.py b/composer/spec/prop_inference.py index 3edc92c0..eb391785 100644 --- a/composer/spec/prop_inference.py +++ b/composer/spec/prop_inference.py @@ -18,6 +18,7 @@ from graphcore.tools.schemas import WithImplementation from composer.input.files import Document +from composer.llm.provider import CacheLevel from composer.spec.context import WorkflowContext, CacheKey, ComponentGroup from composer.spec.graph_builder import bind_standard, run_to_completion from composer.spec.types import PropertyFormulation @@ -358,7 +359,7 @@ async def _run_bug_analysis_inner( "so some of the issues/vulnerabilities/attacks may not be relevant to your analysis. Do *NOT* overfit to this threat model; carefully " "analyze what content of the provided threat model is worth considering vs out of scope. Further, this threat model is just a starting point, " "you should ALSO look for threats *not* mentioned in this document.", - threat_model.to_dict(with_cache=True) + threat_model.to_dict(cache_level=CacheLevel.SHORT) ]) prev_rounds : list[_AgentRoundResult] = [] diff --git a/composer/testing/harness_tape.py b/composer/testing/harness_tape.py index c7eb33e8..7df51e29 100644 --- a/composer/testing/harness_tape.py +++ b/composer/testing/harness_tape.py @@ -206,7 +206,7 @@ def install_fake_llm(fake: Any) -> None: import composer.llm.registry as registry import composer.workflow.services as services - from composer.llm.provider import ProviderServiceBase + from composer.llm.provider import ProviderServiceBase, CacheLevel class _FakeService(ProviderServiceBase): """The real anthropic memory tool (backed by the harness's Postgres @@ -222,7 +222,7 @@ def __init__(self) -> None: class _FakeProvider: provider = _FakeService() - def builder_for(self, *, cache_level: Any = None, disable_thinking: bool = False) -> Any: + def builder_for(self, *, cache_level: CacheLevel = CacheLevel.NONE, disable_thinking: bool = False) -> Any: return _current_fake() fp = _FakeProvider() diff --git a/composer/testing/record_tape.py b/composer/testing/record_tape.py index c7fcf25c..20c4a408 100644 --- a/composer/testing/record_tape.py +++ b/composer/testing/record_tape.py @@ -71,6 +71,7 @@ from langchain_core.outputs import ChatGeneration, LLMResult from composer.diagnostics.timing import get_current_task_id +from composer.llm.types import CacheLevel # task_id used for LLM calls that fire outside any run_task scope. HarnessFakeLLM # raises on such calls, so anything landing here needs manual attention before @@ -208,7 +209,7 @@ def install_recorder(name: str, out_path: str | None = None, *, no_thinking: boo def _wrap_provider(mp: Any) -> Any: orig_builder_for = mp.builder_for - def builder_for(*, cache_level: Any = None, disable_thinking: bool = False) -> Any: + def builder_for(*, cache_level: CacheLevel = CacheLevel.NONE, disable_thinking: bool = False) -> Any: llm = orig_builder_for(cache_level=cache_level, disable_thinking=disable_thinking) if no_thinking: llm = llm.model_copy(update={"thinking": None})