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..49b88aae 100644 --- a/composer/input/files.py +++ b/composer/input/files.py @@ -23,16 +23,18 @@ 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 typing import Protocol, overload -from composer.llm.provider import ProviderKind +from composer.llm.types import CacheLevel + +class ContentRenderer(Protocol): + 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: ... @@ -123,20 +125,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 + 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) @@ -156,33 +152,10 @@ class UploadedFile: basename: str contents: bytes digest: str - provider: ProviderKind - - 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) + renderer: ContentRenderer + + 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 @@ -278,8 +251,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 +272,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 +283,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 +305,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 +323,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 +350,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 +371,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 +385,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..33fe1a30 100644 --- a/composer/llm/anthropic.py +++ b/composer/llm/anthropic.py @@ -5,14 +5,17 @@ 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, + ProviderServiceBase, ProviderSpec ) +from .types import CacheLevel +from .list_iter import ListIter, NoSuchElementError if TYPE_CHECKING: from langchain_core.language_models.chat_models import BaseChatModel @@ -38,7 +41,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() @@ -74,17 +77,54 @@ 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: str, *, cache_level: CacheLevel = CacheLevel.NONE) -> dict: + to_ret : dict[str, Any] = {"type": "text", "text": text} + if (ttl := level_to_ttl(cache_level)) is not None: + to_ret["cache_control"] = { + "type": "ephemeral", + "ttl": ttl + } + return to_ret + + def file_block(self, file_id: str, *, cache_level: CacheLevel = CacheLevel.NONE) -> dict: + to_ret : dict[str, Any] = { + "type": "document", + "source": { + "type": "file", + "file_id": file_id, + }, + } + if (ttl := level_to_ttl(cache_level)) is not None: + to_ret["cache_control"] = { + "type": "ephemeral", + "ttl": ttl + } + 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 +157,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,14 +176,19 @@ 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 + self, *, cache_level: CacheLevel = CacheLevel.NONE, disable_thinking: bool = False ) -> "BaseChatModel": from langchain_anthropic import ChatAnthropic from composer.diagnostics.usage_callback import UsageCallback @@ -153,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( @@ -176,3 +225,8 @@ def builder_for( model_kwargs=model_kwargs, callbacks=[UsageCallback()], ) + +ANTHROPIC_SPEC = ProviderSpec( + matches=matches, + build=AnthropicModelProvider.create +) diff --git a/composer/llm/list_iter.py b/composer/llm/list_iter.py new file mode 100644 index 00000000..b6458fb2 --- /dev/null +++ b/composer/llm/list_iter.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/openai.py b/composer/llm/openai.py index ccca867c..46f2a1fb 100644 --- a/composer/llm/openai.py +++ b/composer/llm/openai.py @@ -14,15 +14,18 @@ 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, ContentRenderer from composer.input.types import ModelConfiguration from composer.llm.provider import ( - ProviderKind, CacheLevel, _ListIter, NoSuchElementError, + ProviderServiceBase, ProviderSpec ) +from .types import CacheLevel +from .list_iter import ListIter, NoSuchElementError if TYPE_CHECKING: from langchain_core.language_models.chat_models import BaseChatModel @@ -67,7 +70,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 +124,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: 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, *, cache_level: CacheLevel = CacheLevel.NONE) -> 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" + 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 @@ -162,6 +185,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,14 +202,14 @@ 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": 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 @@ -207,3 +234,8 @@ def builder_for( max_retries=2, **kwargs, ) + +OPEN_AI_SPEC = ProviderSpec( + matches=matches, + build=OpenAIModelProvider.create +) diff --git a/composer/llm/provider.py b/composer/llm/provider.py index 8967cf75..3814e04d 100644 --- a/composer/llm/provider.py +++ b/composer/llm/provider.py @@ -9,21 +9,48 @@ the per-provider modules can import it without an import cycle. """ -from typing import Literal, Protocol, TYPE_CHECKING -from dataclasses import dataclass, field -import enum +from typing import Protocol, TYPE_CHECKING, Callable +from dataclasses import dataclass +from functools import cached_property +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 - - -type ProviderKind = Literal["anthropic", "openai"] - - -class CacheLevel(enum.StrEnum): - NONE = "none" - SHORT = "short" - LONG = "long" + 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 ModelProvider(Protocol): @@ -33,33 +60,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 + self, *, cache_level: CacheLevel = CacheLevel.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/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/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/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/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..7df51e29 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,10 +206,23 @@ 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, CacheLevel + + 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: + def builder_for(self, *, cache_level: CacheLevel = CacheLevel.NONE, disable_thinking: bool = False) -> Any: return _current_fake() fp = _FakeProvider() @@ -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/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}) 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 f8aac9e9..424b5a70 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -173,3 +173,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) 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