Skip to content

Commit e58a098

Browse files
authored
fix(llm): replace retry with tenacity in ollama.py (apache#367)
## Summary `ollama.py` was using the abandoned `retry` PyPI package (last maintained 2016) while `openai.py` and `litellm.py` in the same directory already use `tenacity`. This PR brings `ollama.py` in line with the established patterns. Closes apache#365
1 parent 5ece55e commit e58a098

4 files changed

Lines changed: 222 additions & 23 deletions

File tree

hugegraph-llm/pyproject.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -47,7 +47,7 @@ dependencies = [
4747
# LLM specific dependencies
4848
"openai",
4949
"ollama",
50-
"retry",
50+
"tenacity",
5151
"tiktoken",
5252
"nltk",
5353
"gradio",

hugegraph-llm/src/hugegraph_llm/models/llms/ollama.py

Lines changed: 31 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -19,8 +19,14 @@
1919
import json
2020
from typing import Any, AsyncGenerator, Callable, Dict, Generator, List, Optional
2121

22+
import httpx
2223
import ollama
23-
from retry import retry
24+
from tenacity import (
25+
retry,
26+
retry_if_exception_type,
27+
stop_after_attempt,
28+
wait_exponential,
29+
)
2430

2531
from hugegraph_llm.models.llms.base import BaseLLM
2632
from hugegraph_llm.utils.log import log
@@ -34,13 +40,17 @@ def __init__(self, model: str, host: str = "127.0.0.1", port: int = 11434, **kwa
3440
self.client = ollama.Client(host=f"http://{host}:{port}", **kwargs)
3541
self.async_client = ollama.AsyncClient(host=f"http://{host}:{port}", **kwargs)
3642

37-
@retry(tries=3, delay=1)
43+
@retry(
44+
stop=stop_after_attempt(3),
45+
wait=wait_exponential(multiplier=1, min=4, max=10),
46+
retry=retry_if_exception_type((ollama.ResponseError, httpx.ConnectError, httpx.TimeoutException)),
47+
)
3848
def generate(
3949
self,
4050
messages: Optional[List[Dict[str, Any]]] = None,
4151
prompt: Optional[str] = None,
4252
) -> str:
43-
"""Comment"""
53+
"""Generate a response to the query messages/prompt."""
4454
if messages is None:
4555
assert prompt is not None, "Messages or prompt must be provided."
4656
messages = [{"role": "user", "content": prompt}]
@@ -56,17 +66,21 @@ def generate(
5666
}
5767
log.info("Token usage: %s", json.dumps(usage))
5868
return response["message"]["content"]
59-
except Exception as e:
60-
print(f"Retrying LLM call {e}")
61-
raise e
62-
63-
@retry(tries=3, delay=1)
69+
except (ollama.ResponseError, httpx.ConnectError, httpx.TimeoutException) as e:
70+
log.error("Retrying LLM call %s", e)
71+
raise
72+
73+
@retry(
74+
stop=stop_after_attempt(3),
75+
wait=wait_exponential(multiplier=1, min=4, max=10),
76+
retry=retry_if_exception_type((ollama.ResponseError, httpx.ConnectError, httpx.TimeoutException)),
77+
)
6478
async def agenerate(
6579
self,
6680
messages: Optional[List[Dict[str, Any]]] = None,
6781
prompt: Optional[str] = None,
6882
) -> str:
69-
"""Comment"""
83+
"""Generate a response to the query messages/prompt asynchronously."""
7084
if messages is None:
7185
assert prompt is not None, "Messages or prompt must be provided."
7286
messages = [{"role": "user", "content": prompt}]
@@ -82,17 +96,17 @@ async def agenerate(
8296
}
8397
log.info("Token usage: %s", json.dumps(usage))
8498
return response["message"]["content"]
85-
except Exception as e:
86-
print(f"Retrying LLM call {e}")
87-
raise e
99+
except (ollama.ResponseError, httpx.ConnectError, httpx.TimeoutException) as e:
100+
log.error("Retrying LLM call %s", e)
101+
raise
88102

89103
def generate_streaming(
90104
self,
91105
messages: Optional[List[Dict[str, Any]]] = None,
92106
prompt: Optional[str] = None,
93107
on_token_callback: Optional[Callable] = None,
94108
) -> Generator[str, None, None]:
95-
"""Comment"""
109+
"""Stream response tokens one by one."""
96110
if messages is None:
97111
assert prompt is not None, "Messages or prompt must be provided."
98112
messages = [{"role": "user", "content": prompt}]
@@ -112,7 +126,7 @@ async def agenerate_streaming(
112126
prompt: Optional[str] = None,
113127
on_token_callback: Optional[Callable] = None,
114128
) -> AsyncGenerator[str, None]:
115-
"""Comment"""
129+
"""Stream response tokens one by one."""
116130
if messages is None:
117131
assert prompt is not None, "Messages or prompt must be provided."
118132
messages = [{"role": "user", "content": prompt}]
@@ -124,9 +138,9 @@ async def agenerate_streaming(
124138
if on_token_callback:
125139
on_token_callback(token)
126140
yield token
127-
except Exception as e:
128-
print(f"Retrying LLM call {e}")
129-
raise e
141+
except (ollama.ResponseError, httpx.ConnectError, httpx.TimeoutException) as e:
142+
log.error("Error in agenerate_streaming: %s", e)
143+
raise
130144

131145
def num_tokens_from_string(
132146
self,

hugegraph-llm/src/tests/models/llms/test_ollama_client.py

Lines changed: 189 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -15,27 +15,212 @@
1515
# specific language governing permissions and limitations
1616
# under the License.
1717

18+
import asyncio
1819
import os
1920
import unittest
21+
from unittest.mock import AsyncMock, MagicMock, patch
22+
23+
import httpx
24+
import ollama
25+
import pytest
26+
from tenacity import RetryError, wait_none
2027

2128
from hugegraph_llm.models.llms.ollama import OllamaClient
2229

30+
pytestmark = pytest.mark.contract
31+
32+
# Minimal dict response matching the structure ollama.Client.chat() returns
33+
_MOCK_RESPONSE = {
34+
"prompt_eval_count": 10,
35+
"eval_count": 5,
36+
"message": {"content": "Paris"},
37+
}
38+
39+
40+
class TestOllamaClientRetryPolicy(unittest.TestCase):
41+
"""Mock-based contract tests for the Tenacity retry policy in OllamaClient.
42+
43+
These tests do not require a running Ollama service. They verify:
44+
- retryable exceptions (ollama.ResponseError, httpx.ConnectError,
45+
httpx.TimeoutException) trigger the configured number of attempts;
46+
- non-retryable exceptions (e.g. ValueError) are NOT retried;
47+
- a transient failure followed by success resolves correctly.
48+
"""
49+
50+
def setUp(self):
51+
# Zero out exponential wait so retry tests complete in milliseconds.
52+
# Tenacity exposes the Retrying object on the decorated function via
53+
# the .retry attribute; its .wait field is mutable.
54+
self._orig_generate_wait = OllamaClient.generate.retry.wait
55+
self._orig_agenerate_wait = OllamaClient.agenerate.retry.wait
56+
OllamaClient.generate.retry.wait = wait_none()
57+
OllamaClient.agenerate.retry.wait = wait_none()
58+
59+
def tearDown(self):
60+
OllamaClient.generate.retry.wait = self._orig_generate_wait
61+
OllamaClient.agenerate.retry.wait = self._orig_agenerate_wait
62+
63+
# ------------------------------------------------------------------ #
64+
# generate() #
65+
# ------------------------------------------------------------------ #
66+
67+
@patch("hugegraph_llm.models.llms.ollama.ollama.Client")
68+
def test_generate_returns_content_on_success(self, mock_client_class):
69+
"""Happy path: generate() returns the message content string."""
70+
mock_client = MagicMock()
71+
mock_client.chat.return_value = _MOCK_RESPONSE
72+
mock_client_class.return_value = mock_client
73+
74+
result = OllamaClient(model="llama3").generate(prompt="hello")
75+
76+
self.assertEqual(result, "Paris")
77+
mock_client.chat.assert_called_once()
78+
79+
@patch("hugegraph_llm.models.llms.ollama.ollama.Client")
80+
def test_generate_retries_response_error_exhausts_all_attempts(self, mock_client_class):
81+
"""ollama.ResponseError is retryable; all 3 attempts are made."""
82+
mock_client = MagicMock()
83+
mock_client.chat.side_effect = ollama.ResponseError("model not found")
84+
mock_client_class.return_value = mock_client
85+
86+
with self.assertRaises(RetryError):
87+
OllamaClient(model="llama3").generate(prompt="hello")
88+
89+
self.assertEqual(mock_client.chat.call_count, 3)
90+
91+
@patch("hugegraph_llm.models.llms.ollama.ollama.Client")
92+
def test_generate_retries_connect_error_exhausts_all_attempts(self, mock_client_class):
93+
"""httpx.ConnectError is retryable; all 3 attempts are made."""
94+
mock_client = MagicMock()
95+
mock_client.chat.side_effect = httpx.ConnectError("connection refused")
96+
mock_client_class.return_value = mock_client
97+
98+
with self.assertRaises(RetryError):
99+
OllamaClient(model="llama3").generate(prompt="hello")
100+
101+
self.assertEqual(mock_client.chat.call_count, 3)
102+
103+
@patch("hugegraph_llm.models.llms.ollama.ollama.Client")
104+
def test_generate_does_not_retry_non_retriable_error(self, mock_client_class):
105+
"""ValueError is not in the retry predicate; only 1 attempt is made."""
106+
mock_client = MagicMock()
107+
mock_client.chat.side_effect = ValueError("unexpected")
108+
mock_client_class.return_value = mock_client
109+
110+
with self.assertRaises(ValueError):
111+
OllamaClient(model="llama3").generate(prompt="hello")
112+
113+
mock_client.chat.assert_called_once()
114+
115+
@patch("hugegraph_llm.models.llms.ollama.ollama.Client")
116+
def test_generate_succeeds_on_second_attempt(self, mock_client_class):
117+
"""Transient ResponseError on attempt 1, success on attempt 2."""
118+
mock_client = MagicMock()
119+
mock_client.chat.side_effect = [
120+
ollama.ResponseError("transient"),
121+
_MOCK_RESPONSE,
122+
]
123+
mock_client_class.return_value = mock_client
124+
125+
result = OllamaClient(model="llama3").generate(prompt="hello")
126+
127+
self.assertEqual(result, "Paris")
128+
self.assertEqual(mock_client.chat.call_count, 2)
129+
130+
# ------------------------------------------------------------------ #
131+
# agenerate() #
132+
# ------------------------------------------------------------------ #
133+
134+
@patch("hugegraph_llm.models.llms.ollama.ollama.AsyncClient")
135+
def test_agenerate_retries_connect_error_exhausts_all_attempts(self, mock_async_client_class):
136+
"""httpx.ConnectError is retryable in agenerate(); all 3 attempts made."""
137+
mock_async_client = MagicMock()
138+
mock_async_client.chat = AsyncMock(side_effect=httpx.ConnectError("connection refused"))
139+
mock_async_client_class.return_value = mock_async_client
140+
141+
async def run():
142+
with self.assertRaises(RetryError):
143+
await OllamaClient(model="llama3").agenerate(prompt="hello")
144+
self.assertEqual(mock_async_client.chat.call_count, 3)
145+
146+
asyncio.run(run())
147+
148+
@patch("hugegraph_llm.models.llms.ollama.ollama.AsyncClient")
149+
def test_agenerate_retries_timeout_exception_exhausts_all_attempts(self, mock_async_client_class):
150+
"""httpx.TimeoutException is retryable in agenerate(); all 3 attempts made."""
151+
mock_async_client = MagicMock()
152+
mock_async_client.chat = AsyncMock(side_effect=httpx.TimeoutException("timed out"))
153+
mock_async_client_class.return_value = mock_async_client
154+
155+
async def run():
156+
with self.assertRaises(RetryError):
157+
await OllamaClient(model="llama3").agenerate(prompt="hello")
158+
self.assertEqual(mock_async_client.chat.call_count, 3)
159+
160+
asyncio.run(run())
161+
162+
@patch("hugegraph_llm.models.llms.ollama.ollama.AsyncClient")
163+
def test_agenerate_does_not_retry_non_retriable_error(self, mock_async_client_class):
164+
"""ValueError is not retryable in agenerate(); only 1 attempt is made."""
165+
mock_async_client = MagicMock()
166+
mock_async_client.chat = AsyncMock(side_effect=ValueError("unexpected"))
167+
mock_async_client_class.return_value = mock_async_client
168+
169+
async def run():
170+
with self.assertRaises(ValueError):
171+
await OllamaClient(model="llama3").agenerate(prompt="hello")
172+
mock_async_client.chat.assert_called_once()
173+
174+
asyncio.run(run())
175+
176+
@patch("hugegraph_llm.models.llms.ollama.ollama.AsyncClient")
177+
def test_agenerate_succeeds_on_second_attempt(self, mock_async_client_class):
178+
"""Transient ResponseError on attempt 1, success on attempt 2."""
179+
mock_async_client = MagicMock()
180+
mock_async_client.chat = AsyncMock(side_effect=[ollama.ResponseError("transient"), _MOCK_RESPONSE])
181+
mock_async_client_class.return_value = mock_async_client
182+
183+
async def run():
184+
result = await OllamaClient(model="llama3").agenerate(prompt="hello")
185+
self.assertEqual(result, "Paris")
186+
self.assertEqual(mock_async_client.chat.call_count, 2)
187+
188+
asyncio.run(run())
189+
190+
191+
class TestOllamaClientExternalService(unittest.TestCase):
192+
"""Integration tests that require a live Ollama service.
193+
194+
Skipped in CI via SKIP_EXTERNAL_SERVICES=true (set in conftest.py).
195+
"""
23196

24-
class TestOllamaClient(unittest.TestCase):
25197
def setUp(self):
26198
self.skip_external = os.getenv("SKIP_EXTERNAL_SERVICES", "false").lower() == "true"
27199

28-
@unittest.skipIf(os.getenv("SKIP_EXTERNAL_SERVICES", "false").lower() == "true", "Skipping external service tests")
200+
@unittest.skipIf(
201+
os.getenv("SKIP_EXTERNAL_SERVICES", "false").lower() == "true",
202+
"Skipping external service tests",
203+
)
29204
def test_generate(self):
30205
ollama_client = OllamaClient(model="llama3:8b-instruct-fp16")
31206
response = ollama_client.generate(prompt="What is the capital of France?")
32207
print(response)
33208

34-
@unittest.skipIf(os.getenv("SKIP_EXTERNAL_SERVICES", "false").lower() == "true", "Skipping external service tests")
209+
@unittest.skipIf(
210+
os.getenv("SKIP_EXTERNAL_SERVICES", "false").lower() == "true",
211+
"Skipping external service tests",
212+
)
35213
def test_stream_generate(self):
36214
ollama_client = OllamaClient(model="llama3:8b-instruct-fp16")
37215

38216
def on_token_callback(chunk):
39217
print(chunk, end="", flush=True)
40218

41-
ollama_client.generate_streaming(prompt="What is the capital of France?", on_token_callback=on_token_callback)
219+
ollama_client.generate_streaming(
220+
prompt="What is the capital of France?",
221+
on_token_callback=on_token_callback,
222+
)
223+
224+
225+
if __name__ == "__main__":
226+
unittest.main()

pyproject.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -114,7 +114,7 @@ constraint-dependencies = [
114114
# LLM dependencies
115115
"openai~=1.61.0",
116116
"ollama~=0.4.8",
117-
"retry~=0.9.2",
117+
"tenacity~=8.5.0",
118118
"tiktoken~=0.7.0",
119119
"nltk~=3.9.1",
120120
"gradio~=5.20.0",

0 commit comments

Comments
 (0)