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
62 changes: 51 additions & 11 deletions sdk/python/src/safety_agent/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,8 @@
import json
import os
import re
from typing import Any
from collections.abc import Awaitable, Callable
from typing import Any, TypeVar

import httpx

Expand Down Expand Up @@ -34,6 +35,9 @@
from .schemas import GUARD_RESPONSE_FORMAT, REDACT_RESPONSE_FORMAT
from .utils.input_processor import process_input, is_vision_model

T = TypeVar("T")
R = TypeVar("R")


def _chunk_text(text: str, chunk_size: int) -> list[str]:
"""
Expand Down Expand Up @@ -105,6 +109,24 @@ def _aggregate_guard_results(results: list[GuardResponse]) -> GuardResponse:
)


async def _map_with_concurrency(
items: list[T],
max_concurrency: int | None,
mapper: Callable[[T], Awaitable[R]],
) -> list[R]:
"""Map items with a per-call limit while preserving input order."""
if not items:
return []

semaphore = asyncio.Semaphore(max_concurrency or len(items))

async def run(item: T) -> R:
async with semaphore:
return await mapper(item)

return list(await asyncio.gather(*(run(item) for item in items)))
Comment on lines +121 to +127


def _parse_json_response(content: str) -> dict[str, Any]:
"""
Parse JSON response, handling both direct JSON and JSON wrapped in markdown code blocks.
Expand Down Expand Up @@ -355,6 +377,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 +398,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: Maximum number of chunk or PDF-page analyses to run concurrently.

Returns:
Response with classification result and token usage
Expand All @@ -387,6 +411,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 +426,19 @@ 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 (
isinstance(max_concurrency, bool)
or not isinstance(max_concurrency, int)
or max_concurrency <= 0
)
):
raise ValueError(
"max_concurrency must be a positive integer, "
f"got {max_concurrency}"
)

# Process the input (handle URLs, bytes, etc.)
processed = await process_input(input)

Expand All @@ -427,11 +465,12 @@ 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 _map_with_concurrency(
non_empty_pages,
max_concurrency,
lambda page_text: self._guard_single_text(
page_text, system_prompt, model, fallback_model
),
)

# Aggregate with OR logic
Expand All @@ -450,11 +489,12 @@ 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 _map_with_concurrency(
chunks,
max_concurrency,
lambda chunk: self._guard_single_text(
chunk, system_prompt, model, fallback_model
),
)

# Aggregate with OR logic
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
"""Maximum number of chunk or PDF-page analyses to run concurrently."""


@dataclass
class GuardClassificationResult:
Expand Down
122 changes: 122 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 All @@ -12,6 +13,7 @@
AnalysisResponse,
AnalysisResponseChoice,
ChatMessage,
ProcessedInput,
TokenUsage,
)

Expand Down Expand Up @@ -229,6 +231,126 @@ async def test_aggregates_token_usage(self, mock_call_provider):
response.usage.prompt_tokens + response.usage.completion_tokens
)

async def test_limits_concurrent_chunk_analysis_and_preserves_result_order(
self, mock_call_provider
):
first_gate = asyncio.Event()
second_gate = asyncio.Event()
started = asyncio.Event()
call_index = 0
in_flight = 0
peak_in_flight = 0

async def provider(*_args, **_kwargs):
nonlocal call_index, in_flight, peak_in_flight
index = call_index
call_index += 1
in_flight += 1
peak_in_flight = max(peak_in_flight, in_flight)
if in_flight == 2:
started.set()

if index == 0:
await first_gate.wait()
elif index == 1:
await second_gate.wait()

in_flight -= 1
return _guard_response("pass", reasoning=f"result-{index}")

mock_call_provider.side_effect = provider
client = create_client(api_key="test-key")
guard_task = asyncio.create_task(
client.guard(
input="Safe content. " * 100,
model="openai/gpt-4o-mini",
chunk_size=50,
max_concurrency=2,
)
)

await asyncio.wait_for(started.wait(), timeout=1)
assert mock_call_provider.call_count == 2
assert peak_in_flight == 2

second_gate.set()
await asyncio.sleep(0)
first_gate.set()

response = await guard_task
assert mock_call_provider.call_count > 2
assert response.reasoning == "result-0"

async def test_limits_concurrent_pdf_page_analysis(self, mock_call_provider):
first_gate = asyncio.Event()
second_gate = asyncio.Event()
started = asyncio.Event()
call_index = 0
in_flight = 0
peak_in_flight = 0

async def provider(*_args, **_kwargs):
nonlocal call_index, in_flight, peak_in_flight
index = call_index
call_index += 1
in_flight += 1
peak_in_flight = max(peak_in_flight, in_flight)
if in_flight == 2:
started.set()

if index == 0:
await first_gate.wait()
elif index == 1:
await second_gate.wait()

in_flight -= 1
return _guard_response("block" if index == 1 else "pass")

mock_call_provider.side_effect = provider
client = create_client(api_key="test-key")

with patch(
"safety_agent.client.process_input", new_callable=AsyncMock
) as mock_process_input:
mock_process_input.return_value = ProcessedInput(
type="pdf",
pages=["Page 1", "Page 2", "Page 3", "Page 4"],
mime_type="application/pdf",
)
guard_task = asyncio.create_task(
client.guard(
input="ignored",
model="openai/gpt-4o-mini",
max_concurrency=2,
)
)

await asyncio.wait_for(started.wait(), timeout=1)
assert mock_call_provider.call_count == 2
assert peak_in_flight == 2

second_gate.set()
await asyncio.sleep(0)
first_gate.set()

response = await guard_task

assert mock_call_provider.call_count == 4
assert response.classification == "block"

@pytest.mark.parametrize("max_concurrency", [0, -1, 1.5, True])
async def test_invalid_max_concurrency_raises(
self, mock_call_provider, max_concurrency
):
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=max_concurrency,
)


class TestGuardErrors:
"""Verify error handling."""
Expand Down
46 changes: 41 additions & 5 deletions sdk/typescript/src/client.ts
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,29 @@ function chunkText(text: string, chunkSize: number): string[] {
return chunks.filter((c) => c.length > 0);
}

async function mapWithConcurrency<T, R>(
items: T[],
maxConcurrency: number | undefined,
mapper: (item: T) => Promise<R>,
): Promise<R[]> {
if (items.length === 0) return [];

const workerCount = Math.min(maxConcurrency ?? items.length, items.length);
const results = new Array<R>(items.length);
let nextIndex = 0;

const worker = async (): Promise<void> => {
while (true) {
const index = nextIndex++;
if (index >= items.length) return;
results[index] = await mapper(items[index]);
}
};

await Promise.all(Array.from({ length: workerCount }, () => 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 +623,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 <= 0)
) {
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 +670,11 @@ export class SafetyClient {
}

// Analyze each page in parallel (similar to chunking strategy)
const results = await Promise.all(
nonEmptyPages.map((pageText) =>
const results = await mapWithConcurrency(
nonEmptyPages,
maxConcurrency,
(pageText) =>
this.guardSingleText(pageText, systemPrompt, model, fallbackModel),
),
);

// Aggregate with OR logic - block if ANY page contains violation
Expand All @@ -661,8 +695,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,
maxConcurrency,
(chunk) => this.guardSingleText(chunk, systemPrompt, model, fallbackModel),
);

// 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;
/** Maximum number of chunk or PDF-page analyses to run concurrently. */
maxConcurrency?: number;
}

/**
Expand Down
Loading