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
56 changes: 39 additions & 17 deletions openagent_eval/providers/llm/ollama.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@

from __future__ import annotations

import time
from typing import Any

import httpx
Expand All @@ -20,7 +21,7 @@
ProviderExecutionError,
)
from openagent_eval.providers.base.llm import LLMProvider
from openagent_eval.providers.models import TokenUsage
from openagent_eval.providers.models import LLMResponse, TokenUsage


class OllamaGenerateRequest(BaseModel):
Expand All @@ -36,7 +37,9 @@ class OllamaGenerateRequest(BaseModel):
model: str = Field(..., description="Model identifier")
prompt: str = Field(..., description="Input prompt")
stream: bool = Field(False, description="Enable streaming")
options: dict[str, Any] = Field(default_factory=dict, description="Generation options")
options: dict[str, Any] = Field(
default_factory=dict, description="Generation options"
)


class OllamaGenerateResponse(BaseModel):
Expand All @@ -58,7 +61,9 @@ class OllamaGenerateResponse(BaseModel):
eval_count: int = Field(0, description="Number of tokens generated")
prompt_eval_count: int = Field(0, description="Number of prompt tokens")
eval_duration: int = Field(0, description="Evaluation duration in nanoseconds")
prompt_eval_duration: int = Field(0, description="Prompt evaluation duration in nanoseconds")
prompt_eval_duration: int = Field(
0, description="Prompt evaluation duration in nanoseconds"
)


class Ollama(LLMProvider):
Expand Down Expand Up @@ -188,7 +193,10 @@ async def generate(self, prompt: str, **kwargs: Any) -> str:
message=f"Ollama API error: {e.response.status_code}",
provider_name=self.name,
original_error=e,
details={"status_code": e.response.status_code, "response": e.response.text},
details={
"status_code": e.response.status_code,
"response": e.response.text,
},
) from e
except Exception as e:
raise ProviderExecutionError(
Expand Down Expand Up @@ -241,25 +249,23 @@ async def get_token_count(self, text: str) -> int:
# Fallback to word-based approximation
return len(text.split())

def generate_with_usage(
self, prompt: str, **kwargs: Any
) -> tuple[str, TokenUsage]:
"""Generate text and return token usage synchronously.
def generate_with_usage(self, prompt: str, **kwargs: Any) -> LLMResponse:
"""Generate text and return a full LLMResponse synchronously.

This is a helper method that generates text and extracts token usage
from Ollama's response metadata. For async usage, call generate()
directly.
from Ollama's response metadata. For async usage, call
generate_with_usage_async() directly.

Args:
prompt: The input prompt.
**kwargs: Optional parameter overrides.

Returns:
Tuple of (generated_text, token_usage).
LLMResponse with content, model, usage, provider, and latency.

Note:
This method is synchronous for convenience. For async contexts,
use generate() directly and track usage separately.
use generate_with_usage_async() instead.
"""
# Build request payload
model = kwargs.get("model", self._model)
Expand All @@ -282,12 +288,14 @@ def generate_with_usage(
timeout=httpx.Timeout(self._timeout),
) as client:
try:
start_time = time.monotonic()
response = client.post(
"/api/generate",
content=request.model_dump_json(),
headers={"Content-Type": "application/json"},
)
response.raise_for_status()
latency_ms = (time.monotonic() - start_time) * 1000
except httpx.ConnectError as e:
raise ProviderConnectionError(
message=f"Failed to connect to Ollama server at {self._base_url}",
Expand Down Expand Up @@ -321,12 +329,18 @@ def generate_with_usage(
total_tokens=total_tokens,
)

return ollama_response.response, usage
return LLMResponse(
content=ollama_response.response,
model=model,
usage=usage,
provider=self.name,
latency_ms=latency_ms,
)

async def generate_with_usage_async(
self, prompt: str, **kwargs: Any
) -> tuple[str, TokenUsage]:
"""Generate text and return token usage asynchronously.
) -> LLMResponse:
"""Generate text and return a full LLMResponse asynchronously.

This method generates text and extracts token usage from Ollama's
response metadata for accurate cost tracking.
Expand All @@ -336,7 +350,7 @@ async def generate_with_usage_async(
**kwargs: Optional parameter overrides.

Returns:
Tuple of (generated_text, token_usage).
LLMResponse with content, model, usage, provider, and latency.
"""
model = kwargs.get("model", self._model)
temperature = kwargs.get("temperature", self._temperature)
Expand All @@ -354,12 +368,14 @@ async def generate_with_usage_async(
)

try:
start_time = time.monotonic()
response = await self._client.post(
"/api/generate",
content=request.model_dump_json(),
headers={"Content-Type": "application/json"},
)
response.raise_for_status()
latency_ms = (time.monotonic() - start_time) * 1000
except httpx.ConnectError as e:
raise ProviderConnectionError(
message=f"Failed to connect to Ollama server at {self._base_url}",
Expand Down Expand Up @@ -393,7 +409,13 @@ async def generate_with_usage_async(
total_tokens=total_tokens,
)

return ollama_response.response, usage
return LLMResponse(
content=ollama_response.response,
model=model,
usage=usage,
provider=self.name,
latency_ms=latency_ms,
)

async def close(self) -> None:
"""Close the HTTP client and clean up resources."""
Expand Down
115 changes: 113 additions & 2 deletions tests/unit/test_providers/test_ollama.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
ProviderExecutionError,
)
from openagent_eval.providers.llm.ollama import Ollama
from openagent_eval.providers.models import LLMResponse, TokenUsage


# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -105,7 +106,9 @@ async def test_generate_success(self, provider: Ollama, mock_httpx_response):
assert result == "Ollama response"

@pytest.mark.asyncio
async def test_generate_with_model_override(self, provider: Ollama, mock_httpx_response):
async def test_generate_with_model_override(
self, provider: Ollama, mock_httpx_response
):
"""generate() respects model override."""
mock_client = AsyncMock()
mock_client.post = AsyncMock(return_value=mock_httpx_response)
Expand Down Expand Up @@ -181,7 +184,9 @@ async def test_generate_parse_error(self, provider: Ollama):
await provider.generate("Test prompt")

@pytest.mark.asyncio
async def test_generate_with_max_tokens(self, provider: Ollama, mock_httpx_response):
async def test_generate_with_max_tokens(
self, provider: Ollama, mock_httpx_response
):
"""generate() passes max_tokens as num_predict in options."""
mock_client = AsyncMock()
mock_client.post = AsyncMock(return_value=mock_httpx_response)
Expand Down Expand Up @@ -233,6 +238,112 @@ async def test_token_count_fallback_on_error(self, provider: Ollama):
assert count == 3 # Three words


# ---------------------------------------------------------------------------
# generate_with_usage() / generate_with_usage_async() tests
#
# Regression for #78: both methods must return an ``LLMResponse`` (matching
# every other LLM provider), not a ``tuple[str, TokenUsage]``.
# ---------------------------------------------------------------------------
class TestOllamaGenerateWithUsage:
"""Tests for the sync/async usage-returning generation helpers."""

def test_generate_with_usage_returns_llm_response(
self, provider: Ollama, mock_httpx_response, monkeypatch
):
"""generate_with_usage() returns an LLMResponse with all fields."""
import openagent_eval.providers.llm.ollama as ollama_module

mock_client = MagicMock()
mock_client.post = MagicMock(return_value=mock_httpx_response)
client_cm = MagicMock()
client_cm.__enter__ = MagicMock(return_value=mock_client)
client_cm.__exit__ = MagicMock(return_value=False)
monkeypatch.setattr(
ollama_module.httpx, "Client", MagicMock(return_value=client_cm)
)

result = provider.generate_with_usage("Test prompt")

assert isinstance(result, LLMResponse)
assert result.content == "Ollama response"
assert result.model == "llama3.2"
assert result.provider == "ollama"
assert isinstance(result.usage, TokenUsage)
assert result.usage.prompt_tokens == 10
assert result.usage.completion_tokens == 15
assert result.usage.total_tokens == 25
assert result.latency_ms >= 0.0

def test_generate_with_usage_model_override(
self, provider: Ollama, mock_httpx_response, monkeypatch
):
"""generate_with_usage() reflects a per-call model override."""
import openagent_eval.providers.llm.ollama as ollama_module

mock_client = MagicMock()
mock_client.post = MagicMock(return_value=mock_httpx_response)
client_cm = MagicMock()
client_cm.__enter__ = MagicMock(return_value=mock_client)
client_cm.__exit__ = MagicMock(return_value=False)
monkeypatch.setattr(
ollama_module.httpx, "Client", MagicMock(return_value=client_cm)
)

result = provider.generate_with_usage("Test prompt", model="mistral")

assert isinstance(result, LLMResponse)
assert result.model == "mistral"

@pytest.mark.asyncio
async def test_generate_with_usage_async_returns_llm_response(
self, provider: Ollama, mock_httpx_response
):
"""generate_with_usage_async() returns an LLMResponse with all fields."""
mock_client = AsyncMock()
mock_client.post = AsyncMock(return_value=mock_httpx_response)
provider._client = mock_client

result = await provider.generate_with_usage_async("Test prompt")

assert isinstance(result, LLMResponse)
assert result.content == "Ollama response"
assert result.model == "llama3.2"
assert result.provider == "ollama"
assert isinstance(result.usage, TokenUsage)
assert result.usage.prompt_tokens == 10
assert result.usage.completion_tokens == 15
assert result.usage.total_tokens == 25
assert result.latency_ms >= 0.0

@pytest.mark.asyncio
async def test_generate_with_usage_async_model_override(
self, provider: Ollama, mock_httpx_response
):
"""generate_with_usage_async() reflects a per-call model override."""
mock_client = AsyncMock()
mock_client.post = AsyncMock(return_value=mock_httpx_response)
provider._client = mock_client

result = await provider.generate_with_usage_async(
"Test prompt", model="mistral"
)

assert isinstance(result, LLMResponse)
assert result.model == "mistral"

@pytest.mark.asyncio
async def test_generate_with_usage_async_connection_error(self, provider: Ollama):
"""generate_with_usage_async() raises ProviderConnectionError on ConnectError."""
import httpx

mock_client = AsyncMock()
mock_client.post = AsyncMock(side_effect=httpx.ConnectError("refused"))
provider._client = mock_client

with pytest.raises(ProviderConnectionError, match="connect"):
await provider.generate_with_usage_async("Test prompt")


# ---------------------------------------------------------------------------
# Context manager tests
# ---------------------------------------------------------------------------
Expand Down
Loading