diff --git a/src/bub/builtin/model_runner.py b/src/bub/builtin/model_runner.py index 79f5829a..bd7365a6 100644 --- a/src/bub/builtin/model_runner.py +++ b/src/bub/builtin/model_runner.py @@ -11,6 +11,7 @@ from any_llm import AnyLLM from any_llm.constants import LLMProvider +from any_llm.providers.anthropic.base import BaseAnthropicProvider from any_llm.providers.openai.base import BaseOpenAIProvider from any_llm.types.completion import ( ChatCompletion, @@ -46,17 +47,13 @@ CompletionResult = ChatCompletion | ParsedChatCompletion[Any] | AsyncIterator[ChatCompletionChunk] -def _stream_usage_options(llm: AnyLLM, *, stream: bool) -> dict[str, Any] | None: - """Make streaming completions report token usage. - - OpenAI-style streaming responses omit the `usage` block unless the request - sets `stream_options.include_usage`; without it every streamed run records - zero tokens (and zero cost). Only OpenAI-compatible providers accept the - field, so gate on the provider base class — anthropic/gemini reject it. - """ - if stream and isinstance(llm, BaseOpenAIProvider): - return {"include_usage": True} - return None +def _extra_options(llm: AnyLLM, *, stream: bool) -> dict[str, Any]: + """Return provider-specific extra completion options.""" + if isinstance(llm, BaseAnthropicProvider): + return {"cache_control": {"type": "ephemeral"}} + elif stream and isinstance(llm, BaseOpenAIProvider): + return {"stream_options": {"include_usage": True}} + return {} class ModelRunner: @@ -101,7 +98,7 @@ async def completion_response( tools=tool_payloads, max_tokens=max_tokens if max_tokens is not None else self.settings.max_tokens, stream=streaming, - stream_options=_stream_usage_options(llm, stream=streaming), + **_extra_options(llm, stream=streaming), ) except Exception as exc: if completion_error is None: diff --git a/tests/test_builtin_model_runner.py b/tests/test_builtin_model_runner.py index 7daebe71..f60e4543 100644 --- a/tests/test_builtin_model_runner.py +++ b/tests/test_builtin_model_runner.py @@ -6,6 +6,7 @@ import pytest from any_llm.constants import LLMProvider +from any_llm.providers.anthropic.base import BaseAnthropicProvider from any_llm.providers.openai.base import BaseOpenAIProvider from any_llm.types.completion import ChatCompletionChunk @@ -53,6 +54,23 @@ async def stream() -> AsyncIterator[ChatCompletionChunk]: return stream() +class _FakeStreamingAnthropicProvider(BaseAnthropicProvider): + def __init__(self) -> None: + self.completion_kwargs: dict[str, Any] | None = None + + def _init_client(self, api_key: str | None = None, api_base: str | None = None, **kwargs: Any) -> None: + pass + + async def acompletion(self, **kwargs: Any) -> AsyncIterator[ChatCompletionChunk]: + self.completion_kwargs = kwargs + + async def stream() -> AsyncIterator[ChatCompletionChunk]: + if False: + yield + + return stream() + + class _FakeOpenAIModelRunner(ModelRunner): def __init__(self, settings: AgentSettings, llm: _FakeStreamingOpenAIProvider) -> None: super().__init__(settings) @@ -62,6 +80,15 @@ def iter_llm_clients(self, model: str) -> Iterator[tuple[ModelCandidate, _FakeSt yield ModelCandidate(provider=LLMProvider.OPENAI, model_id=model, name=f"openai:{model}"), self._llm +class _FakeAnthropicModelRunner(ModelRunner): + def __init__(self, settings: AgentSettings, llm: _FakeStreamingAnthropicProvider) -> None: + super().__init__(settings) + self._llm = llm + + def iter_llm_clients(self, model: str) -> Iterator[tuple[ModelCandidate, _FakeStreamingAnthropicProvider]]: + yield ModelCandidate(provider=LLMProvider.ANTHROPIC, model_id=model, name=f"anthropic:{model}"), self._llm + + @pytest.mark.asyncio async def test_streaming_openai_usage_is_requested_and_recorded_in_tape(tmp_path: Path) -> None: store = InMemoryTapeStore() @@ -93,3 +120,19 @@ async def test_streaming_openai_usage_is_requested_and_recorded_in_tape(tmp_path "prompt_tokens": 3, "total_tokens": 5, } + + +@pytest.mark.asyncio +async def test_anthropic_prompt_caching_is_requested() -> None: + llm = _FakeStreamingAnthropicProvider() + runner = _FakeAnthropicModelRunner( + AgentSettings.model_construct(model="anthropic:claude-test", max_tokens=100), + llm, + ) + + await runner.completion_response(model="claude-test", messages=[{"role": "user", "content": "hello"}], tools=[]) + + assert llm.completion_kwargs is not None + assert llm.completion_kwargs["stream"] is True + assert llm.completion_kwargs["cache_control"] == {"type": "ephemeral"} + assert "stream_options" not in llm.completion_kwargs diff --git a/uv.lock b/uv.lock index 709d7f14..ada2f4ad 100644 --- a/uv.lock +++ b/uv.lock @@ -167,7 +167,7 @@ wheels = [ [[package]] name = "any-llm-sdk" -version = "1.17.0" +version = "1.22.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "anthropic" }, @@ -178,9 +178,9 @@ dependencies = [ { name = "rich" }, { name = "typing-extensions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/31/19/1fd535059dbb55b68251586be8d8aa7253f3abb3b4d4aa4a9fa9784334da/any_llm_sdk-1.17.0.tar.gz", hash = "sha256:429d59eac71e2dcdeff3848cf55a5ae71dfde012496f0c932d6c302b6ea11ca1", size = 149396, upload-time = "2026-06-05T09:34:00.451Z" } +sdist = { url = "https://files.pythonhosted.org/packages/bf/de/a29738e1c8702a79a514f9521430fe77686fb9f9122980f4899a7baf58b5/any_llm_sdk-1.22.1.tar.gz", hash = "sha256:b07c6fcef6d7fc13e0f0419bc1464f9316b17e663c34382fb7722a906241f635", size = 164931, upload-time = "2026-07-22T15:43:29.12Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/42/ed/e2891f42247892e3358db9327bf2bb3fb7a0f046c4fb891c4d22dbdc26a1/any_llm_sdk-1.17.0-py3-none-any.whl", hash = "sha256:634d1e1413bdd2c593aa8eec0d32c4cad868c681b31da4f78e36e239833e9c76", size = 198742, upload-time = "2026-06-05T09:33:59.052Z" }, + { url = "https://files.pythonhosted.org/packages/03/5a/85c558cc9fb6fd8a7fd3e7523c9c20545316190d22372026ddbf53cdc0c9/any_llm_sdk-1.22.1-py3-none-any.whl", hash = "sha256:950339b68ca1d99fb123c523d246140e1c947b0ff26e3fb9bc37194a49f3b509", size = 217209, upload-time = "2026-07-22T15:43:27.664Z" }, ] [[package]]