diff --git a/pydantic_ai_slim/pydantic_ai/concurrency.py b/pydantic_ai_slim/pydantic_ai/concurrency.py index ec5611852f..05a16c5231 100644 --- a/pydantic_ai_slim/pydantic_ai/concurrency.py +++ b/pydantic_ai_slim/pydantic_ai/concurrency.py @@ -12,6 +12,8 @@ from opentelemetry.trace import Tracer, get_tracer from typing_extensions import Self +from .exceptions import UserError + __all__ = ( 'AbstractConcurrencyLimiter', 'ConcurrencyLimiter', @@ -20,6 +22,11 @@ ) +def _validate_max_running(max_running: int) -> None: + if max_running < 1: + raise UserError(f'max_running must be >= 1, got {max_running}. Use None for no concurrency limiting.') + + class AbstractConcurrencyLimiter(ABC): """Abstract base class for concurrency limiters. @@ -67,7 +74,7 @@ class ConcurrencyLimit: """Configuration for concurrency limiting with optional backpressure. Args: - max_running: Maximum number of concurrent operations allowed. + max_running: Maximum number of concurrent operations allowed. Must be >= 1. max_queued: Maximum number of operations waiting in the queue. If None, the queue is unlimited. If exceeded, raises `ConcurrencyLimitExceeded`. """ @@ -75,6 +82,9 @@ class ConcurrencyLimit: max_running: int max_queued: int | None = None + def __post_init__(self) -> None: + _validate_max_running(self.max_running) + class ConcurrencyLimiter(AbstractConcurrencyLimiter): """A concurrency limiter that tracks waiting operations for observability. @@ -95,12 +105,16 @@ def __init__( """Initialize the ConcurrencyLimiter. Args: - max_running: Maximum number of concurrent operations. + max_running: Maximum number of concurrent operations. Must be >= 1. max_queued: Maximum queue depth before raising ConcurrencyLimitExceeded. name: Optional name for this limiter, used for observability when sharing a limiter across multiple models or agents. tracer: OpenTelemetry tracer for span creation. + + Raises: + UserError: If `max_running` is less than 1. """ + _validate_max_running(max_running) self._limiter = anyio.CapacityLimiter(max_running) self._max_queued = max_queued self._name = name diff --git a/tests/test_concurrency.py b/tests/test_concurrency.py index b0317f1a47..89d46aae69 100644 --- a/tests/test_concurrency.py +++ b/tests/test_concurrency.py @@ -10,7 +10,8 @@ import pytest from pydantic_ai import Agent, ConcurrencyLimit, ConcurrencyLimiter, ConcurrencyLimitExceeded -from pydantic_ai.concurrency import get_concurrency_context +from pydantic_ai.concurrency import get_concurrency_context, normalize_to_limiter +from pydantic_ai.exceptions import UserError from pydantic_ai.models.test import TestModel if TYPE_CHECKING: @@ -66,6 +67,18 @@ async def test_nowait_acquisition(self): async with get_concurrency_context(limiter, 'test'): pass # No waiting + @pytest.mark.parametrize('max_running', [0, -1, -5]) + async def test_invalid_max_running(self, max_running: int): + """Test that invalid concurrency limits are rejected before creating the limiter.""" + with pytest.raises(UserError, match=f'max_running must be >= 1, got {max_running}'): + ConcurrencyLimiter(max_running=max_running) + + @pytest.mark.parametrize('max_running', [0, -1, -5]) + async def test_invalid_concurrency_limit_config(self, max_running: int): + """Test that invalid concurrency limit config values fail eagerly.""" + with pytest.raises(UserError, match=f'max_running must be >= 1, got {max_running}'): + ConcurrencyLimit(max_running=max_running) + async def test_waiting_count_tracking(self): """Test that waiting_count is accurately tracked.""" limiter = ConcurrencyLimiter(max_running=1) @@ -193,6 +206,22 @@ async def test_from_limiter_config(self): assert limiter.max_running == 5 assert limiter._max_queued == 10 + async def test_from_invalid_limit(self): + """Test that invalid normalized concurrency limits fail eagerly.""" + with pytest.raises(UserError, match='max_running must be >= 1, got 0'): + ConcurrencyLimiter.from_limit(0) + + async def test_normalize_to_limiter_rejects_invalid_limit(self): + """Test that invalid normalized concurrency limits fail before use.""" + with pytest.raises(UserError, match='max_running must be >= 1, got 0'): + normalize_to_limiter(0) + + async def test_none_still_means_no_limit(self): + """None must remain the sanctioned 'no limiting' path after rejecting non-positive limits.""" + assert normalize_to_limiter(None) is None + agent = Agent(TestModel(), max_concurrency=None) + assert agent._concurrency_limiter is None + async def test_properties(self): """Test the various properties of ConcurrencyLimiter.""" limiter = ConcurrencyLimiter(max_running=5, name='test-limiter') @@ -304,6 +333,12 @@ async def test_agent_with_limiter_concurrency(self): assert agent._concurrency_limiter.max_running == 5 assert agent._concurrency_limiter._max_queued == 10 + @pytest.mark.parametrize('max_running', [0, -1]) + async def test_agent_invalid_int_concurrency_limit(self, max_running: int): + """Test that invalid agent concurrency limits fail eagerly at construction.""" + with pytest.raises(UserError, match=f'max_running must be >= 1, got {max_running}'): + Agent(TestModel(), max_concurrency=max_running) + class TestConcurrencyLimitedModel: """Tests for the ConcurrencyLimitedModel wrapper."""