Skip to content
This repository was archived by the owner on Jul 10, 2026. It is now read-only.

Commit f7c5040

Browse files
fix(rate-limiter): prevent deepcopy from creating duplicate instances (#89)
* fix(rate-limiter): prevent deepcopy from creating duplicate instances Implement singleton pattern for RateLimiter to ensure shared instances are preserved during deepcopy operations. This fixes a bug where rate limiters were being incorrectly duplicated when generators were copied or serialized. - Make RateLimiter frozen to prevent mutation - Add model_validator to enforce singleton pattern - Override __deepcopy__ to return same instance - Add tests to verify rate limiter preservation in with_params and serialization * chore: code formatting * chore: improve comments Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> * fix(test): unique rate limiter name * chore: code formatting * fix(test): update expectation * fix(test): update expectation --------- Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
1 parent e3d8825 commit f7c5040

3 files changed

Lines changed: 89 additions & 7 deletions

File tree

src/giskard/agents/rate_limiter.py

Lines changed: 24 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -3,15 +3,15 @@
33
import uuid
44
from contextlib import asynccontextmanager
55

6-
from pydantic import BaseModel, Field, PrivateAttr
6+
from pydantic import BaseModel, Field, PrivateAttr, model_validator
77

88

99
class RateLimiterStrategy(BaseModel):
1010
min_interval: float
1111
max_concurrent: int = Field(default=5)
1212

1313

14-
class RateLimiter(BaseModel):
14+
class RateLimiter(BaseModel, frozen=True):
1515
rate_limiter_id: str = Field(default_factory=lambda: str(uuid.uuid4()))
1616
strategy: RateLimiterStrategy
1717

@@ -80,6 +80,28 @@ async def acquire(self) -> None:
8080
def release(self) -> None:
8181
self._semaphore.release()
8282

83+
@model_validator(mode="wrap")
84+
@classmethod
85+
def _validate_singleton(cls, v, handler, _info) -> "RateLimiter":
86+
# Determine the rate limiter ID from the input
87+
rate_limiter_id = None
88+
if isinstance(v, dict):
89+
rate_limiter_id = v.get("rate_limiter_id")
90+
elif isinstance(v, RateLimiter):
91+
rate_limiter_id = v.rate_limiter_id
92+
93+
# If ID exists in our global registry, return the singleton
94+
if rate_limiter_id and rate_limiter_id in _rate_limiters:
95+
return _rate_limiters[rate_limiter_id]
96+
97+
# 3. Otherwise, proceed with standard validation/creation
98+
instance = handler(v)
99+
return instance
100+
101+
def __deepcopy__(self, memo) -> "RateLimiter":
102+
# RateLimiter is a shared resource, so we can just return the same instance.
103+
return self
104+
83105

84106
_rate_limiters: dict[str, RateLimiter] = {}
85107

tests/test_generator.py

Lines changed: 63 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
from giskard.agents.chat import Chat, Message
66
from giskard.agents.generators.base import GenerationParams, Response
77
from giskard.agents.generators.litellm_generator import LiteLLMGenerator
8-
from giskard.agents.rate_limiter import RateLimiter
8+
from giskard.agents.rate_limiter import RateLimiter, _rate_limiters
99
from giskard.agents.templates import MessageTemplate
1010
from giskard.agents.workflow import ChatWorkflow
1111
from litellm import ModelResponse
@@ -119,8 +119,12 @@ async def test_generator_without_rate_limiter(mock_response):
119119

120120

121121
async def test_generator_rate_limiter_context():
122-
rate_limiter = RateLimiter.from_rpm(rpm=100, rate_limiter_id="test")
123-
generator = LiteLLMGenerator(model="test-model", rate_limiter="test")
122+
rate_limiter = RateLimiter.from_rpm(
123+
rpm=100, rate_limiter_id="test_generator_rate_limiter_context"
124+
)
125+
generator = LiteLLMGenerator(
126+
model="test-model", rate_limiter="test_generator_rate_limiter_context"
127+
)
124128
assert generator.rate_limiter is rate_limiter
125129

126130

@@ -139,6 +143,62 @@ def test_generator_with_params():
139143
assert generator.params.response_format is None
140144

141145

146+
def test_generator_with_params_and_rate_limiter():
147+
"""Test that with_params works correctly with a rate limiter."""
148+
rate_limiter = RateLimiter.from_rpm(rpm=100, max_concurrent=5)
149+
generator = LiteLLMGenerator(model="test-model", rate_limiter=rate_limiter)
150+
151+
# Verify initial state
152+
assert generator.rate_limiter is rate_limiter
153+
154+
# Call with_params and verify rate limiter is preserved
155+
generator_with_params = generator.with_params(temperature=0.5, max_tokens=100)
156+
assert generator_with_params.params.temperature == 0.5
157+
assert generator_with_params.params.max_tokens == 100
158+
# Verify rate limiter is preserved and the same instance
159+
assert generator_with_params.rate_limiter is rate_limiter
160+
161+
# Verify original generator is unchanged
162+
assert generator.params.temperature == 1.0 # default value
163+
assert generator.params.max_tokens is None
164+
assert generator.rate_limiter is rate_limiter
165+
166+
167+
def test_generator_serialization_keep_rate_limiter_instance():
168+
"""Test that serializing and deserializing a generator preserves the rate limiter instance."""
169+
rate_limiter = RateLimiter.from_rpm(rpm=100, max_concurrent=5)
170+
generator = LiteLLMGenerator(model="test-model", rate_limiter=rate_limiter)
171+
172+
json_str = generator.model_dump_json()
173+
deserialized_generator = LiteLLMGenerator.model_validate_json(json_str)
174+
175+
assert deserialized_generator.rate_limiter is rate_limiter
176+
177+
178+
def test_generator_serialization_recreate_rate_limiter_instance_if_not_in_registry():
179+
"""Test that deserializing a generator recreates the rate limiter if it's not in the registry."""
180+
rate_limiter = RateLimiter.from_rpm(rpm=100, max_concurrent=5)
181+
generator = LiteLLMGenerator(model="test-model", rate_limiter=rate_limiter)
182+
183+
json_str = generator.model_dump_json()
184+
del _rate_limiters[rate_limiter.rate_limiter_id]
185+
deserialized_generator = LiteLLMGenerator.model_validate_json(json_str)
186+
187+
assert deserialized_generator.rate_limiter is not rate_limiter
188+
assert (
189+
deserialized_generator.rate_limiter.rate_limiter_id
190+
== rate_limiter.rate_limiter_id
191+
)
192+
assert (
193+
deserialized_generator.rate_limiter.strategy.min_interval
194+
== rate_limiter.strategy.min_interval
195+
)
196+
assert (
197+
deserialized_generator.rate_limiter.strategy.max_concurrent
198+
== rate_limiter.strategy.max_concurrent
199+
)
200+
201+
142202
async def test_generator_with_params_overwrite(mock_response):
143203
# ARRANGE: Create a generator with base parameters.
144204
generator = LiteLLMGenerator(model="test-model").with_params(

tests/test_rate_limiter.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -121,9 +121,9 @@ async def test_rate_limiter_max_concurrent():
121121

122122
def test_can_serialize_rate_limiter():
123123
rate_limiter = RateLimiter.from_rpm(
124-
rpm=600, max_concurrent=10, rate_limiter_id="test"
124+
rpm=600, max_concurrent=10, rate_limiter_id="test_can_serialize_rate_limiter"
125125
)
126126
assert rate_limiter.model_dump() == {
127-
"rate_limiter_id": "test",
127+
"rate_limiter_id": "test_can_serialize_rate_limiter",
128128
"strategy": {"min_interval": 0.1, "max_concurrent": 10},
129129
}

0 commit comments

Comments
 (0)