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
106 changes: 65 additions & 41 deletions src/any_llm/providers/groq/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
from groq.types import ModelListResponse as GroqModelListResponse
from groq.types.chat import ChatCompletion as GroqChatCompletion
from groq.types.chat import ChatCompletionChunk as GroqChatCompletionChunk
from groq.types.completion_usage import CompletionUsage as GroqCompletionUsage

from any_llm.types.completion import (
ChatCompletion,
Expand Down Expand Up @@ -35,6 +36,17 @@
)


def _extract_groq_timing_details(usage: GroqCompletionUsage) -> dict[str, float]:
"""Return Groq timing fields in seconds, omitting the ones the provider did not report."""
timing = {
"queue_time": usage.queue_time,
"prompt_time": usage.prompt_time,
"completion_time": usage.completion_time,
"total_time": usage.total_time,
}
return {field: value for field, value in timing.items() if value is not None}


def to_chat_completion(response: GroqChatCompletion) -> ChatCompletion:
"""Convert Groq ChatCompletion into our ChatCompletion type directly."""

Expand All @@ -44,11 +56,15 @@ def to_chat_completion(response: GroqChatCompletion) -> ChatCompletion:
# Reference: https://console.groq.com/docs/prompt-caching
prompt_details = response.usage.prompt_tokens_details
cached_tokens = prompt_details.cached_tokens if prompt_details else None
usage = CompletionUsage(
prompt_tokens=response.usage.prompt_tokens,
completion_tokens=response.usage.completion_tokens,
total_tokens=response.usage.total_tokens,
prompt_tokens_details=PromptTokensDetails(cached_tokens=cached_tokens) if cached_tokens else None,
timing_details = _extract_groq_timing_details(response.usage)
usage = CompletionUsage.model_validate(
{
"prompt_tokens": response.usage.prompt_tokens,
"completion_tokens": response.usage.completion_tokens,
"total_tokens": response.usage.total_tokens,
"prompt_tokens_details": PromptTokensDetails(cached_tokens=cached_tokens) if cached_tokens else None,
**timing_details,
}
)

choices: list[Choice] = []
Expand Down Expand Up @@ -108,57 +124,65 @@ def to_chat_completion(response: GroqChatCompletion) -> ChatCompletion:
def _create_openai_chunk_from_groq_chunk(groq_chunk: GroqChatCompletionChunk) -> ChatCompletionChunk:
"""Convert a Groq streaming chunk to OpenAI ChatCompletionChunk format."""

choice_data = groq_chunk.choices[0]
delta_data = choice_data.delta
choices: list[ChunkChoice] = []
if groq_chunk.choices:
choice_data = groq_chunk.choices[0]
delta_data = choice_data.delta

delta = ChoiceDelta(
content=delta_data.content,
reasoning=Reasoning(content=delta_data.reasoning) if delta_data.reasoning else None,
role=cast("Literal['developer', 'system', 'user', 'assistant', 'tool'] | None", delta_data.role),
)
delta = ChoiceDelta(
content=delta_data.content,
reasoning=Reasoning(content=delta_data.reasoning) if delta_data.reasoning else None,
role=cast("Literal['developer', 'system', 'user', 'assistant', 'tool'] | None", delta_data.role),
)

if delta_data.tool_calls:
openai_tool_calls = []
for tool_call in delta_data.tool_calls:
openai_tool_call = ChoiceDeltaToolCall(
index=tool_call.index if tool_call.index is not None else 0,
id=tool_call.id,
type="function",
function=ChoiceDeltaToolCallFunction(
name=tool_call.function.name if tool_call.function else None,
arguments=tool_call.function.arguments if tool_call.function else None,
if delta_data.tool_calls:
openai_tool_calls = []
for tool_call in delta_data.tool_calls:
openai_tool_call = ChoiceDeltaToolCall(
index=tool_call.index if tool_call.index is not None else 0,
id=tool_call.id,
type="function",
function=ChoiceDeltaToolCallFunction(
name=tool_call.function.name if tool_call.function else None,
arguments=tool_call.function.arguments if tool_call.function else None,
)
if tool_call.function
else None,
)
if tool_call.function
else None,
openai_tool_calls.append(openai_tool_call)
delta.tool_calls = openai_tool_calls
else:
delta.tool_calls = None

choices.append(
ChunkChoice(
index=choice_data.index,
delta=delta,
finish_reason=choice_data.finish_reason,
)
openai_tool_calls.append(openai_tool_call)
delta.tool_calls = openai_tool_calls
else:
delta.tool_calls = None

choice = ChunkChoice(
index=choice_data.index,
delta=delta,
finish_reason=choice_data.finish_reason,
)
)

usage = None
usage_data = groq_chunk.usage
usage_data = groq_chunk.usage or (groq_chunk.x_groq.usage if groq_chunk.x_groq else None)
if usage_data:
# Groq's prompt_tokens already includes cached tokens (cached_tokens is a subset).
# Reference: https://console.groq.com/docs/prompt-caching
prompt_details = usage_data.prompt_tokens_details
cached_tokens = prompt_details.cached_tokens if prompt_details else None
usage = CompletionUsage(
prompt_tokens=usage_data.prompt_tokens,
completion_tokens=usage_data.completion_tokens,
total_tokens=usage_data.total_tokens,
prompt_tokens_details=PromptTokensDetails(cached_tokens=cached_tokens) if cached_tokens else None,
timing_details = _extract_groq_timing_details(usage_data)
usage = CompletionUsage.model_validate(
{
"prompt_tokens": usage_data.prompt_tokens,
"completion_tokens": usage_data.completion_tokens,
"total_tokens": usage_data.total_tokens,
"prompt_tokens_details": PromptTokensDetails(cached_tokens=cached_tokens) if cached_tokens else None,
**timing_details,
}
)

return ChatCompletionChunk(
id=groq_chunk.id,
choices=[choice],
choices=choices,
created=groq_chunk.created,
model=groq_chunk.model,
object="chat.completion.chunk",
Expand Down
37 changes: 28 additions & 9 deletions src/any_llm/providers/ollama/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,17 @@ def _create_openai_embedding_response_from_ollama(
)


def _extract_ollama_timing_details(response: OllamaChatResponse) -> dict[str, int]:
"""Return Ollama timing fields in nanoseconds, omitting the ones the provider did not report."""
timing = {
"total_duration": response.total_duration,
"load_duration": response.load_duration,
"prompt_eval_duration": response.prompt_eval_duration,
"eval_duration": response.eval_duration,
}
return {field: value for field, value in timing.items() if value is not None}


def _create_openai_chunk_from_ollama_chunk(ollama_chunk: OllamaChatResponse) -> ChatCompletionChunk:
"""Convert an Ollama streaming chunk to OpenAI ChatCompletionChunk format."""

Expand Down Expand Up @@ -125,11 +136,15 @@ def _create_openai_chunk_from_ollama_chunk(ollama_chunk: OllamaChatResponse) ->
usage = None
prompt_tokens = ollama_chunk.prompt_eval_count
completion_tokens = ollama_chunk.eval_count
if prompt_tokens or completion_tokens:
usage = CompletionUsage(
prompt_tokens=prompt_tokens or 0,
completion_tokens=completion_tokens or 0,
total_tokens=(prompt_tokens or 0) + (completion_tokens or 0),
timing_details = _extract_ollama_timing_details(ollama_chunk)
if prompt_tokens or completion_tokens or timing_details:
usage = CompletionUsage.model_validate(
{
"prompt_tokens": prompt_tokens or 0,
"completion_tokens": completion_tokens or 0,
"total_tokens": (prompt_tokens or 0) + (completion_tokens or 0),
**timing_details,
}
)

return ChatCompletionChunk(
Expand Down Expand Up @@ -158,6 +173,7 @@ def _create_chat_completion_from_ollama_response(response: OllamaChatResponse) -

prompt_tokens = response.prompt_eval_count or 0
completion_tokens = response.eval_count or 0
timing_details = _extract_ollama_timing_details(response)

response_message: OllamaMessage = response.message
if not response_message or not isinstance(response_message, OllamaMessage):
Expand Down Expand Up @@ -202,10 +218,13 @@ def _create_chat_completion_from_ollama_response(response: OllamaChatResponse) -

choice = Choice(index=0, finish_reason=finish_reason, message=message)

usage = CompletionUsage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
usage = CompletionUsage.model_validate(
{
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
**timing_details,
}
)

return ChatCompletion(
Expand Down
106 changes: 106 additions & 0 deletions tests/integration/test_provider_timing.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,106 @@
"""Integration tests for issue #1258: provider-reported timing must survive normalization.

Groq reports timing in seconds on its usage object, Ollama in nanoseconds on its chat
response. Both land in ``usage.model_extra``. The unit tests hand-build the provider SDK
objects, so only these tests prove the fields are actually on the wire, which matters most
for Groq streaming: groq 1.6.0 takes no ``stream_options``, so streaming usage arrives under
``x_groq.usage`` rather than top-level ``chunk.usage``.

Requires GROQ_API_KEY for the Groq tests and a reachable Ollama host for the Ollama ones.
"""

from collections.abc import AsyncIterator
from typing import Any

import httpx
import pytest

from any_llm import AnyLLM, LLMProvider
from any_llm.exceptions import MissingApiKeyError
from any_llm.types.completion import ChatCompletion, ChatCompletionChunk, ChatCompletionMessage

_GROQ_TIMING_FIELDS = ("queue_time", "prompt_time", "completion_time", "total_time")
_OLLAMA_TIMING_FIELDS = ("total_duration", "load_duration", "prompt_eval_duration", "eval_duration")

_PROMPT: list[dict[str, Any] | ChatCompletionMessage] = [{"role": "user", "content": "Reply with the single word OK."}]


@pytest.mark.asyncio
async def test_groq_timing_non_streaming(provider_model_map: dict[LLMProvider, str]) -> None:
"""Groq's per-request timing survives on a non-streaming completion."""
try:
llm = AnyLLM.create(LLMProvider.GROQ)
except MissingApiKeyError:
pytest.skip("Groq API key not provided, skipping")

result = await llm.acompletion(model=provider_model_map[LLMProvider.GROQ], messages=_PROMPT)

assert isinstance(result, ChatCompletion)
assert result.usage is not None
extras = result.usage.model_extra or {}
assert set(_GROQ_TIMING_FIELDS) <= extras.keys(), f"missing Groq timing fields in {extras}"
assert extras["total_time"] > 0


@pytest.mark.asyncio
async def test_groq_timing_streaming(provider_model_map: dict[LLMProvider, str]) -> None:
"""Groq's streaming usage and timing arrive on the final chunk via x_groq."""
try:
llm = AnyLLM.create(LLMProvider.GROQ)
except MissingApiKeyError:
pytest.skip("Groq API key not provided, skipping")

stream = await llm.acompletion(model=provider_model_map[LLMProvider.GROQ], messages=_PROMPT, stream=True)
assert isinstance(stream, AsyncIterator)

usages = []
async for chunk in stream:
assert isinstance(chunk, ChatCompletionChunk)
if chunk.usage is not None:
usages.append(chunk.usage)

assert usages, "no chunk reported usage: streaming usage is not reaching the converter"
extras = usages[-1].model_extra or {}
assert set(_GROQ_TIMING_FIELDS) <= extras.keys(), f"missing Groq timing fields in {extras}"
assert usages[-1].completion_tokens > 0


@pytest.mark.asyncio
async def test_ollama_timing_non_streaming(provider_model_map: dict[LLMProvider, str]) -> None:
"""Ollama's duration fields survive on a non-streaming completion."""
llm = AnyLLM.create(LLMProvider.OLLAMA)

try:
result = await llm.acompletion(model=provider_model_map[LLMProvider.OLLAMA], messages=_PROMPT)
# An unreachable host surfaces as a builtin ConnectionError from the ollama SDK on the
# non-streaming call and as a raw httpx.ConnectError on the streaming one.
except (ConnectionError, httpx.ConnectError, httpx.HTTPStatusError):
pytest.skip("Local Ollama host is not set up, skipping")
Comment on lines +73 to +78

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🎯 Functional Correctness | 🟠 Major | ⚡ Quick win

Do not skip all HTTP response failures.

httpx.HTTPStatusError includes invalid model names, authentication failures, and provider failures. These conditions must fail the test. Skip only confirmed unavailable-service conditions, such as connection errors or connection timeouts. Include the concrete connection failure in the skip message.

  • tests/integration/test_provider_timing.py#L73-L78: Remove httpx.HTTPStatusError from the skip handler.
  • tests/integration/test_provider_timing.py#L92-L102: Remove httpx.HTTPStatusError from the skip handler.

As per coding guidelines, “Every test skip and fix must have a concrete root cause”.

📍 Affects 1 file
  • tests/integration/test_provider_timing.py#L73-L78 (this comment)
  • tests/integration/test_provider_timing.py#L92-L102
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@tests/integration/test_provider_timing.py` around lines 73 - 78, Update the
exception handlers around both provider timing calls in
tests/integration/test_provider_timing.py (73-78 and 92-102) to stop skipping on
httpx.HTTPStatusError, while continuing to skip only confirmed connection
failures or timeouts. Include the caught connection exception details in each
pytest.skip message so every skip identifies its concrete root cause.

Source: Coding guidelines


assert isinstance(result, ChatCompletion)
assert result.usage is not None
extras = result.usage.model_extra or {}
assert set(_OLLAMA_TIMING_FIELDS) <= extras.keys(), f"missing Ollama timing fields in {extras}"
assert extras["total_duration"] > 0


@pytest.mark.asyncio
async def test_ollama_timing_streaming(provider_model_map: dict[LLMProvider, str]) -> None:
"""Ollama reports its duration fields on the final streaming chunk."""
llm = AnyLLM.create(LLMProvider.OLLAMA)

try:
stream = await llm.acompletion(model=provider_model_map[LLMProvider.OLLAMA], messages=_PROMPT, stream=True)
assert isinstance(stream, AsyncIterator)

usages = []
async for chunk in stream:
assert isinstance(chunk, ChatCompletionChunk)
if chunk.usage is not None:
usages.append(chunk.usage)
except (ConnectionError, httpx.ConnectError, httpx.HTTPStatusError):
pytest.skip("Local Ollama host is not set up, skipping")

assert usages, "no chunk reported usage"
extras = usages[-1].model_extra or {}
assert set(_OLLAMA_TIMING_FIELDS) <= extras.keys(), f"missing Ollama timing fields in {extras}"
Loading
Loading