Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions composer/cli/console_codegen.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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():
Expand Down
4 changes: 2 additions & 2 deletions composer/cli/tui_codegen.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)

Expand Down
7 changes: 2 additions & 5 deletions composer/cli/tui_pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 (
Expand Down
75 changes: 23 additions & 52 deletions composer/input/files.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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: ...


Expand Down Expand Up @@ -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)
Expand All @@ -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
Expand Down Expand Up @@ -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: ...
Expand All @@ -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
Expand All @@ -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(
Expand All @@ -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(
Expand All @@ -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(
Expand All @@ -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(
Expand All @@ -400,20 +371,20 @@ 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
renderable ``Document``: text stays inline as ``InMemoryTextFile``;
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)

86 changes: 70 additions & 16 deletions composer/llm/anthropic.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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()
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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(
Expand All @@ -176,3 +225,8 @@ def builder_for(
model_kwargs=model_kwargs,
callbacks=[UsageCallback()],
)

ANTHROPIC_SPEC = ProviderSpec(
matches=matches,
build=AnthropicModelProvider.create
)
Loading
Loading