diff --git a/sdk/python/src/safety_agent/client.py b/sdk/python/src/safety_agent/client.py index 8cf7dfb22..58e424973 100644 --- a/sdk/python/src/safety_agent/client.py +++ b/sdk/python/src/safety_agent/client.py @@ -6,7 +6,7 @@ import json import os import re -from typing import Any +from typing import Any, Awaitable, Callable, TypeVar import httpx @@ -70,6 +70,30 @@ def _chunk_text(text: str, chunk_size: int) -> list[str]: return chunks +T = TypeVar("T") +R = TypeVar("R") + + +async def _gather_with_concurrency( + items: list[T], + fn: Callable[[T], Awaitable[R]], + max_concurrency: int | None = None, +) -> list[R]: + """Runs `fn` over `items` with at most `max_concurrency` in flight, preserving order. + Unset or >= len(items) falls back to plain `asyncio.gather`. + """ + if not max_concurrency or max_concurrency >= len(items): + return list(await asyncio.gather(*[fn(item) for item in items])) + + semaphore = asyncio.Semaphore(max_concurrency) + + async def bounded(item: T) -> R: + async with semaphore: + return await fn(item) + + return list(await asyncio.gather(*[bounded(item) for item in items])) + + def _aggregate_guard_results(results: list[GuardResponse]) -> GuardResponse: """ Aggregate multiple guard results using OR logic. @@ -355,6 +379,7 @@ async def guard( fallback_model: str | None = None, system_prompt: str | None = None, chunk_size: int = 8000, + max_concurrency: int | None = None, # Also accept GuardOptions-style kwargs **kwargs: Any, ) -> GuardResponse: @@ -375,6 +400,7 @@ async def guard( fallback_model: Fallback model when the primary returns a retryable error (429/500/502/503) system_prompt: Optional custom system prompt chunk_size: Characters per chunk. Default: 8000. Set to 0 to disable chunking. + max_concurrency: Max concurrent provider requests for multi-chunk/PDF-page analysis. Default: unbounded. Returns: Response with classification result and token usage @@ -387,6 +413,7 @@ async def guard( fallback_model = fallback_model or options.fallback_model system_prompt = system_prompt or options.system_prompt chunk_size = options.chunk_size + max_concurrency = options.max_concurrency # Handle input passed via kwargs if input is None: @@ -401,6 +428,11 @@ async def guard( if chunk_size < 0: raise ValueError(f"chunk_size must be non-negative, got {chunk_size}") + if max_concurrency is not None and max_concurrency < 1: + raise ValueError( + f"max_concurrency must be a positive integer, got {max_concurrency}" + ) + # Process the input (handle URLs, bytes, etc.) processed = await process_input(input) @@ -427,15 +459,16 @@ async def guard( ) # Analyze each page in parallel - results = await asyncio.gather( - *[ - self._guard_single_text(page_text, system_prompt, model, fallback_model) - for page_text in non_empty_pages - ] + results = await _gather_with_concurrency( + non_empty_pages, + lambda page_text: self._guard_single_text( + page_text, system_prompt, model, fallback_model + ), + max_concurrency, ) # Aggregate with OR logic - aggregated = _aggregate_guard_results(list(results)) + aggregated = _aggregate_guard_results(results) self._post_usage(aggregated.usage) return aggregated @@ -450,15 +483,14 @@ async def guard( # Chunk and process in parallel chunks = _chunk_text(text, chunk_size) - results = await asyncio.gather( - *[ - self._guard_single_text(chunk, system_prompt, model, fallback_model) - for chunk in chunks - ] + results = await _gather_with_concurrency( + chunks, + lambda chunk: self._guard_single_text(chunk, system_prompt, model, fallback_model), + max_concurrency, ) # Aggregate with OR logic - aggregated = _aggregate_guard_results(list(results)) + aggregated = _aggregate_guard_results(results) self._post_usage(aggregated.usage) return aggregated diff --git a/sdk/python/src/safety_agent/types.py b/sdk/python/src/safety_agent/types.py index 41eecfbaf..6c339d216 100644 --- a/sdk/python/src/safety_agent/types.py +++ b/sdk/python/src/safety_agent/types.py @@ -989,6 +989,9 @@ class GuardOptions: chunk_size: int = 8000 """Characters per chunk. Default: 8000. Set to 0 to disable chunking.""" + max_concurrency: int | None = None + """Max concurrent provider requests when analyzing multiple chunks or PDF pages. Default: unbounded.""" + @dataclass class GuardClassificationResult: diff --git a/sdk/python/tests/test_guard.py b/sdk/python/tests/test_guard.py index 9954aa66b..5d6a1e318 100644 --- a/sdk/python/tests/test_guard.py +++ b/sdk/python/tests/test_guard.py @@ -2,6 +2,7 @@ Guard unit tests – mocks call_provider to avoid hitting real APIs. """ +import asyncio import json from unittest.mock import AsyncMock, patch @@ -230,6 +231,66 @@ async def test_aggregates_token_usage(self, mock_call_provider): ) +class TestGuardMaxConcurrency: + """Verify max_concurrency bounds in-flight provider calls.""" + + async def test_caps_in_flight_calls(self, mock_call_provider): + in_flight = 0 + max_in_flight = 0 + + async def tracked(*_args, **_kwargs): + nonlocal in_flight, max_in_flight + in_flight += 1 + max_in_flight = max(max_in_flight, in_flight) + await asyncio.sleep(0.005) + in_flight -= 1 + return _guard_response("pass") + + mock_call_provider.side_effect = tracked + client = create_client(api_key="test-key") + + await client.guard( + input="A " * 5000, + model="openai/gpt-4o-mini", + chunk_size=100, + max_concurrency=3, + ) + + assert mock_call_provider.call_count > 3 + assert max_in_flight <= 3 + + async def test_full_concurrency_when_unset(self, mock_call_provider): + in_flight = 0 + max_in_flight = 0 + + async def tracked(*_args, **_kwargs): + nonlocal in_flight, max_in_flight + in_flight += 1 + max_in_flight = max(max_in_flight, in_flight) + await asyncio.sleep(0.005) + in_flight -= 1 + return _guard_response("pass") + + mock_call_provider.side_effect = tracked + client = create_client(api_key="test-key") + + await client.guard( + input="A " * 5000, + model="openai/gpt-4o-mini", + chunk_size=100, + ) + + assert max_in_flight == mock_call_provider.call_count + + async def test_non_positive_max_concurrency_raises(self, mock_call_provider): + client = create_client(api_key="test-key") + + with pytest.raises(ValueError, match="max_concurrency must be a positive integer"): + await client.guard( + input="Test", model="openai/gpt-4o-mini", max_concurrency=0 + ) + + class TestGuardErrors: """Verify error handling.""" diff --git a/sdk/typescript/src/client.ts b/sdk/typescript/src/client.ts index da364f7eb..5f13f3cf1 100644 --- a/sdk/typescript/src/client.ts +++ b/sdk/typescript/src/client.ts @@ -61,6 +61,33 @@ function chunkText(text: string, chunkSize: number): string[] { return chunks.filter((c) => c.length > 0); } +/** + * Runs `fn` over `items` with at most `maxConcurrency` in flight, preserving order. + * Unset or >= items.length falls back to plain `Promise.all`. + */ +async function mapWithConcurrency( + items: T[], + fn: (item: T) => Promise, + maxConcurrency?: number, +): Promise { + if (!maxConcurrency || maxConcurrency >= items.length) { + return Promise.all(items.map(fn)); + } + + const results: R[] = new Array(items.length); + let nextIndex = 0; + + async function worker() { + while (nextIndex < items.length) { + const i = nextIndex++; + results[i] = await fn(items[i]); + } + } + + await Promise.all(Array.from({ length: maxConcurrency }, worker)); + return results; +} + /** * Aggregate multiple guard results using OR logic * Block if ANY chunk is blocked, merge all violations @@ -600,6 +627,7 @@ export class SafetyClient { model = DEFAULT_GUARD_MODEL, fallbackModel, chunkSize = 8000, + maxConcurrency, } = options; // Validate chunkSize is non-negative @@ -607,6 +635,15 @@ export class SafetyClient { throw new Error(`chunkSize must be non-negative, got ${chunkSize}`); } + if ( + maxConcurrency !== undefined && + (!Number.isInteger(maxConcurrency) || maxConcurrency < 1) + ) { + throw new Error( + `maxConcurrency must be a positive integer, got ${maxConcurrency}`, + ); + } + // Process the input (handle URLs, Blobs, etc.) const processed = await processInput(input); @@ -637,10 +674,10 @@ export class SafetyClient { } // Analyze each page in parallel (similar to chunking strategy) - const results = await Promise.all( - nonEmptyPages.map((pageText) => - this.guardSingleText(pageText, systemPrompt, model, fallbackModel), - ), + const results = await mapWithConcurrency( + nonEmptyPages, + (pageText) => this.guardSingleText(pageText, systemPrompt, model, fallbackModel), + maxConcurrency, ); // Aggregate with OR logic - block if ANY page contains violation @@ -661,8 +698,10 @@ export class SafetyClient { // Chunk and process in parallel const chunks = chunkText(text, chunkSize); - const results = await Promise.all( - chunks.map((chunk) => this.guardSingleText(chunk, systemPrompt, model, fallbackModel)), + const results = await mapWithConcurrency( + chunks, + (chunk) => this.guardSingleText(chunk, systemPrompt, model, fallbackModel), + maxConcurrency, ); // Aggregate with OR logic diff --git a/sdk/typescript/src/types.ts b/sdk/typescript/src/types.ts index f5586c97b..ce8053b25 100644 --- a/sdk/typescript/src/types.ts +++ b/sdk/typescript/src/types.ts @@ -1085,6 +1085,8 @@ export interface GuardOptions { fallbackModel?: SupportedModel; /** Characters per chunk. Default: 8000. Set to 0 to disable chunking. */ chunkSize?: number; + /** Max concurrent provider requests when analyzing multiple chunks or PDF pages. Default: unbounded. */ + maxConcurrency?: number; } /** diff --git a/sdk/typescript/tests/chunking.test.ts b/sdk/typescript/tests/chunking.test.ts index 76ea359d3..2cdd507fe 100644 --- a/sdk/typescript/tests/chunking.test.ts +++ b/sdk/typescript/tests/chunking.test.ts @@ -218,4 +218,76 @@ describe("Guard Chunking", () => { expect(response.cwe_codes).toContain("CWE-79"); }); }); + + describe("maxConcurrency", () => { + it("should cap in-flight provider calls at maxConcurrency", async () => { + const largeInput = "A ".repeat(5000); + let inFlight = 0; + let maxInFlight = 0; + mockCallProvider.mockImplementation(async () => { + inFlight++; + maxInFlight = Math.max(maxInFlight, inFlight); + await new Promise((resolve) => setTimeout(resolve, 5)); + inFlight--; + return guardResponse("pass"); + }); + + const client = createClient({ apiKey: "test-key" }); + await client.guard({ + input: largeInput, + model: "openai/gpt-4o-mini", + chunkSize: 100, + maxConcurrency: 3, + }); + + expect(mockCallProvider.mock.calls.length).toBeGreaterThan(3); + expect(maxInFlight).toBeLessThanOrEqual(3); + }); + + it("should allow full concurrency when maxConcurrency is not set", async () => { + const largeInput = "A ".repeat(5000); + let inFlight = 0; + let maxInFlight = 0; + mockCallProvider.mockImplementation(async () => { + inFlight++; + maxInFlight = Math.max(maxInFlight, inFlight); + await new Promise((resolve) => setTimeout(resolve, 5)); + inFlight--; + return guardResponse("pass"); + }); + + const client = createClient({ apiKey: "test-key" }); + await client.guard({ + input: largeInput, + model: "openai/gpt-4o-mini", + chunkSize: 100, + }); + + expect(maxInFlight).toBe(mockCallProvider.mock.calls.length); + }); + + it("should throw for a non-positive maxConcurrency", async () => { + const client = createClient({ apiKey: "test-key" }); + + await expect( + client.guard({ + input: "What's the weather like today?", + model: "openai/gpt-4o-mini", + maxConcurrency: 0, + }), + ).rejects.toThrow("maxConcurrency must be a positive integer"); + }); + + it("should throw for a non-integer maxConcurrency", async () => { + const client = createClient({ apiKey: "test-key" }); + + await expect( + client.guard({ + input: "What's the weather like today?", + model: "openai/gpt-4o-mini", + maxConcurrency: 1.5, + }), + ).rejects.toThrow("maxConcurrency must be a positive integer"); + }); + }); });