Skip to content
Open
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
58 changes: 45 additions & 13 deletions sdk/python/src/safety_agent/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
import json
import os
import re
from typing import Any
from typing import Any, Awaitable, Callable, TypeVar

import httpx

Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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:
Expand All @@ -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
Expand All @@ -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:
Expand All @@ -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)

Expand All @@ -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

Expand All @@ -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

Expand Down
3 changes: 3 additions & 0 deletions sdk/python/src/safety_agent/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
61 changes: 61 additions & 0 deletions sdk/python/tests/test_guard.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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."""

Expand Down
51 changes: 45 additions & 6 deletions sdk/typescript/src/client.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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<T, R>(
items: T[],
fn: (item: T) => Promise<R>,
maxConcurrency?: number,
): Promise<R[]> {
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
Expand Down Expand Up @@ -600,13 +627,23 @@ export class SafetyClient {
model = DEFAULT_GUARD_MODEL,
fallbackModel,
chunkSize = 8000,
maxConcurrency,
} = options;

// Validate chunkSize is non-negative
if (chunkSize < 0) {
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);

Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
2 changes: 2 additions & 0 deletions sdk/typescript/src/types.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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;
}

/**
Expand Down
72 changes: 72 additions & 0 deletions sdk/typescript/tests/chunking.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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");
});
});
});