diff --git a/packages/sie_gateway/openapi.json b/packages/sie_gateway/openapi.json index 15375fab..68b1e554 100644 --- a/packages/sie_gateway/openapi.json +++ b/packages/sie_gateway/openapi.json @@ -1518,6 +1518,18 @@ "grammar": { "description": "Optional grammar object accepted by the gateway grammar validator." }, + "images": { + "description": "Optional inline images paired with ``prompt``. When present, the worker\nrenders one user turn through the model's native chat template.", + "items": { + "$ref": "#/components/schemas/NativeGenerateImage" + }, + "maxItems": 16, + "minItems": 1, + "type": [ + "array", + "null" + ] + }, "logit_bias": { "additionalProperties": { "format": "double", @@ -2268,6 +2280,31 @@ ], "type": "object" }, + "NativeGenerateImage": { + "description": "One inline image on the SIE-native generate surface.", + "properties": { + "data": { + "description": "Canonical standard-base64 encoded image bytes, at most 16 MiB decoded.\nRemote URLs are not accepted.", + "maxLength": 22369624, + "minLength": 1, + "type": "string" + }, + "format": { + "description": "Short media-format hint such as ``png`` or ``jpeg``.", + "maxLength": 32, + "minLength": 1, + "pattern": "^[A-Za-z0-9.+-]+$", + "type": [ + "string", + "null" + ] + } + }, + "required": [ + "data" + ], + "type": "object" + }, "OpenAIEmbeddingDataEntry": { "properties": { "embedding": { @@ -5853,7 +5890,7 @@ }, "/v1/generate/{model}": { "post": { - "description": "SIE-native text generation. Omit `stream` or set it to false for a blocking JSON response; set `stream: true` for SIE-native Server-Sent Events terminated by `data: [DONE]`. The model path parameter must use the SIE-safe ID (for example `Qwen__Qwen3-4B-Instruct`); HF-style slashes reject with 400.", + "description": "SIE-native generation with optional bounded inline images. Omit `stream` or set it to false for a blocking JSON response; set `stream: true` for SIE-native Server-Sent Events terminated by `data: [DONE]`. The model path parameter must use the SIE-safe ID (for example `Qwen__Qwen3-4B-Instruct`); HF-style slashes reject with 400.", "operationId": "proxy_generate", "parameters": [ { diff --git a/packages/sie_gateway/src/handlers/proxy.rs b/packages/sie_gateway/src/handlers/proxy.rs index 7138dc3a..9fd2ca86 100644 --- a/packages/sie_gateway/src/handlers/proxy.rs +++ b/packages/sie_gateway/src/handlers/proxy.rs @@ -35,7 +35,13 @@ use super::models::{extract_bearer_token, mask_token}; const GATEWAY_VERSION: &str = env!("CARGO_PKG_VERSION"); const MAX_PROXY_BODY: usize = 16 * 1024 * 1024; -const MAX_GENERATE_BODY: usize = 4 * 1024 * 1024; +// Native generate accepts bounded inline images. A 16 MiB decoded image grows +// to ~21.4 MiB in base64; leave room for the prompt and JSON envelope while +// still bounding aggregate request memory at ingress. +const MAX_GENERATE_BODY: usize = 24 * 1024 * 1024; +const MAX_GENERATE_IMAGES: usize = 16; +const MAX_GENERATE_IMAGE_BYTES: usize = 16 * 1024 * 1024; +const MAX_GENERATE_IMAGE_BASE64_CHARS: usize = 4 * MAX_GENERATE_IMAGE_BYTES.div_ceil(3); // The audio preprocessor accepts 24 MiB of encoded media. Base64 expands that // to exactly 32 MiB; leave bounded room for the surrounding native JSON item. const MAX_EXTRACT_BODY: usize = 34 * 1024 * 1024; @@ -1666,10 +1672,10 @@ pub(crate) async fn proxy_request( .unwrap_or(false); let use_msgpack_out = publisher::wants_msgpack(req.headers()); - // Per-endpoint body cap. `generate` is pure text — Qwen3.5's 32k - // context is ~128 KiB of UTF-8, so 4 MiB gives ~30× headroom and - // closes the trivial-OOM-under-concurrency vector that the legacy - // 256 MiB cap left open. The remaining endpoints routed here + // Per-endpoint body cap. `generate` allows one 16 MiB decoded inline + // image after JSON base64 expansion, so its 24 MiB cap closes the + // trivial-OOM-under-concurrency vector while admitting the maximum + // legal image. The remaining endpoints routed here // (encode / score) get the same text-appropriate 16 MiB cap as the // chat / embeddings paths. Extract accepts bounded binary media, so // its cap covers the maximum legal audio after JSON base64 expansion. @@ -1971,6 +1977,28 @@ async fn queue_mode_proxy( } } } + let has_images = params + .generate + .as_ref() + .is_some_and(generate_params_have_images); + if has_images { + let image_supported = state + .model_registry + .get_model_info(model) + .is_some_and(|info| info.info_extras.supports_vision_generation()); + if !native_generate_image_input_allowed(params.generate.as_ref(), image_supported) { + return ( + StatusCode::BAD_REQUEST, + Json(json_openai_error( + format!("Model '{display_model}' does not support image input"), + oai_type::INVALID_REQUEST, + Some("images"), + oai_code::UNSUPPORTED_FIELD, + )), + ) + .into_response(); + } + } } // Generate has its own publish + result-collection path @@ -9659,6 +9687,7 @@ fn parse_lora_adapter_field(value: Option<&serde_json::Value>) -> Result, format?: }`` object and remains +/// base64 on the generate work wire so it can round-trip through the +/// sidecar's ``serde_json::Value`` representation. +#[allow(clippy::result_large_err)] +fn parse_native_generate_images( + value: Option<&serde_json::Value>, +) -> Result>, Response> { + let Some(value) = value else { + return Ok(None); + }; + if value.is_null() { + return Ok(None); + } + let Some(entries) = value.as_array() else { + return Err(sampler_bad_request( + "'images' must be a non-empty array".to_string(), + "images", + oai_code::INVALID_REQUEST, + )); + }; + if entries.is_empty() || entries.len() > MAX_GENERATE_IMAGES { + return Err(sampler_bad_request( + format!("'images' must contain between 1 and {MAX_GENERATE_IMAGES} entries"), + "images", + oai_code::INVALID_REQUEST, + )); + } + let mut images = Vec::with_capacity(entries.len()); + for (index, entry) in entries.iter().enumerate() { + let param = format!("images[{index}]"); + let Some(object) = entry.as_object() else { + return Err(sampler_bad_request( + format!("'{param}' must be an object"), + ¶m, + oai_code::INVALID_REQUEST, + )); + }; + if let Some(unknown) = object + .keys() + .find(|key| !matches!(key.as_str(), "data" | "format")) + { + let unknown_param = format!("{param}.{unknown}"); + return Err(sampler_bad_request( + format!("'{unknown_param}' is not supported"), + &unknown_param, + oai_code::UNSUPPORTED_FIELD, + )); + } + let data_param = format!("{param}.data"); + let Some(data) = object.get("data").and_then(serde_json::Value::as_str) else { + return Err(sampler_bad_request( + format!("'{data_param}' must be a non-empty base64 string"), + &data_param, + oai_code::INVALID_REQUEST, + )); + }; + if !native_generate_image_encoded_size_allowed(data.len()) { + return Err(sampler_bad_request( + format!( + "'{data_param}' must decode to between 1 byte and {MAX_GENERATE_IMAGE_BYTES} bytes" + ), + &data_param, + oai_code::INVALID_REQUEST, + )); + } + let decoded = base64::engine::general_purpose::STANDARD + .decode(data) + .map_err(|_| { + sampler_bad_request( + format!("'{data_param}' must be valid standard base64"), + &data_param, + oai_code::INVALID_REQUEST, + ) + })?; + if decoded.is_empty() || decoded.len() > MAX_GENERATE_IMAGE_BYTES { + return Err(sampler_bad_request( + format!( + "'{data_param}' must decode to between 1 byte and {MAX_GENERATE_IMAGE_BYTES} bytes" + ), + &data_param, + oai_code::INVALID_REQUEST, + )); + } + let format = match object.get("format") { + None | Some(serde_json::Value::Null) => None, + Some(serde_json::Value::String(format)) + if !format.is_empty() + && format.len() <= 32 + && format.bytes().all(|byte| { + byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'+' | b'-') + }) => + { + Some(format.to_ascii_lowercase()) + } + Some(_) => { + let format_param = format!("{param}.format"); + return Err(sampler_bad_request( + format!("'{format_param}' must be a short media-format token"), + &format_param, + oai_code::INVALID_REQUEST, + )); + } + }; + images.push(publisher::ChatImage { + data: data.to_string(), + format, + }); + } + Ok(Some(images)) +} + +fn native_generate_image_encoded_size_allowed(length: usize) -> bool { + (1..=MAX_GENERATE_IMAGE_BASE64_CHARS).contains(&length) +} + +fn generate_params_have_images(params: &publisher::GenerateParams) -> bool { + matches!( + ¶ms.input, + publisher::GenerateInput::Messages { messages } + if messages.iter().any(|message| { + message.images.as_ref().is_some_and(|images| !images.is_empty()) + }) + ) +} + +fn native_generate_image_input_allowed( + params: Option<&publisher::GenerateParams>, + model_supports_vision_generation: bool, +) -> bool { + !params.is_some_and(generate_params_have_images) || model_supports_vision_generation +} + /// Parse the walking-skeleton ``/v1/generate/{model}`` JSON body shape into a /// :class:`GenerateParams`. Returns ``Ok(None)`` when required fields are /// missing or malformed — the caller surfaces a generic 400 in that case. @@ -9685,9 +9849,9 @@ const GENERATE_ACCEPTED_FIELDS: &[&str] = &[ /// that response is the precise OpenAI envelope and is returned /// verbatim. /// -/// Only the ``Prompt`` arm is exposed via this entrypoint; chat requests -/// flow through :func:`chat_params_from_json` and assemble -/// the ``Messages`` arm there. +/// Text-only requests preserve the ``Prompt`` arm. Native image requests and +/// chat requests assemble the existing ``Messages`` arm so both reach the +/// worker's model-native template renderer. #[allow(clippy::result_large_err)] fn generate_params_from_json( parsed: &serde_json::Value, @@ -9704,6 +9868,7 @@ fn generate_params_from_json( )) } }; + let images = parse_native_generate_images(parsed.get("images"))?; // Granular `max_new_tokens` rejection so OpenAI-shaped error // envelopes carry the correct `param` field. Previously every // failure mode collapsed to `Ok(None)` and the caller emitted @@ -9911,8 +10076,21 @@ fn generate_params_from_json( let top_logprobs = parse_top_logprobs_field(parsed.get("top_logprobs"))?; check_logprobs_consistency(logprobs, top_logprobs)?; + let input = match images { + Some(images) => publisher::GenerateInput::Messages { + messages: vec![publisher::ChatMessage { + role: "user".to_string(), + content: prompt, + tool_calls: None, + tool_call_id: None, + images: Some(images), + content_parts: None, + }], + }, + None => publisher::GenerateInput::Prompt { prompt }, + }; Ok(Some(publisher::GenerateParams { - input: publisher::GenerateInput::Prompt { prompt }, + input, max_new_tokens, temperature, top_p, @@ -10060,6 +10238,11 @@ fn generate_params_from_rmpv( )) } }; + let images = parse_native_generate_images( + rmpv_map_get(parsed, "images") + .map(rmpv_to_json_owned) + .as_ref(), + )?; // Mirror the JSON path's granular error attribution so SDKs that // branch on `error.param` see the same field name whether the // wire format is JSON or msgpack. @@ -10311,8 +10494,21 @@ fn generate_params_from_rmpv( let top_logprobs = parse_top_logprobs_field(bridge("top_logprobs").as_ref())?; check_logprobs_consistency(logprobs, top_logprobs)?; + let input = match images { + Some(images) => publisher::GenerateInput::Messages { + messages: vec![publisher::ChatMessage { + role: "user".to_string(), + content: prompt, + tool_calls: None, + tool_call_id: None, + images: Some(images), + content_parts: None, + }], + }, + None => publisher::GenerateInput::Prompt { prompt }, + }; Ok(Some(publisher::GenerateParams { - input: publisher::GenerateInput::Prompt { prompt }, + input, max_new_tokens, temperature, top_p, @@ -13163,6 +13359,131 @@ mod tests { assert!(params.grammar.is_none()); } + #[test] + fn test_generate_params_from_json_maps_native_images_to_worker_messages() { + for stream in [false, true] { + let body = serde_json::json!({ + "prompt": "Read the image", + "images": [{"data": "aGVsbG8=", "format": "PNG"}], + "max_new_tokens": 8, + "stream": stream, + }); + let params = _expect_generate_ok(&body); + assert_eq!(params.stream, stream); + assert!(generate_params_have_images(¶ms)); + let publisher::GenerateInput::Messages { messages } = params.input else { + panic!("image-native generate must use the worker message path") + }; + assert_eq!(messages.len(), 1); + assert_eq!(messages[0].role, "user"); + assert_eq!(messages[0].content, "Read the image"); + let images = messages[0].images.as_ref().expect("one image"); + assert_eq!(images[0].data, "aGVsbG8="); + assert_eq!(images[0].format.as_deref(), Some("png")); + } + } + + #[tokio::test] + async fn test_generate_params_from_json_rejects_malformed_native_images() { + let too_many = serde_json::Value::Array(vec![ + serde_json::json!({"data": "aGk="}); + MAX_GENERATE_IMAGES + 1 + ]); + let cases = [ + (serde_json::json!([]), "images"), + (serde_json::json!([{"data": "!!!"}]), "images[0].data"), + (serde_json::json!([{"data": "aGk"}]), "images[0].data"), + (serde_json::json!([{"data": "__8="}]), "images[0].data"), + (serde_json::json!([{"data": "AB=="}]), "images[0].data"), + ( + serde_json::json!([{"data": "aGVsbG8=", "url": "https://example.com/image.png"}]), + "images[0].url", + ), + ( + serde_json::json!([{"data": "aGVsbG8=", "format": "png;bad"}]), + "images[0].format", + ), + (too_many, "images"), + ]; + for (images, expected_param) in cases { + let body = serde_json::json!({ + "prompt": "Read", + "images": images, + "max_new_tokens": 8, + }); + let error = _expect_generate_err(&body).await; + assert_eq!(error["error"]["param"], expected_param); + } + } + + #[test] + fn test_generate_params_from_rmpv_maps_native_images_to_worker_messages() { + let body = serde_json::json!({ + "prompt": "Read", + "images": [{"data": "aGVsbG8=", "format": "png"}], + "max_new_tokens": 8, + }); + let rmpv::Value::Map(map) = json_to_rmpv(body) else { + panic!("object must map") + }; + let params = generate_params_from_rmpv(&map) + .expect("valid request") + .expect("generate params"); + assert!(generate_params_have_images(¶ms)); + } + + #[test] + fn test_generate_native_image_validation_has_json_msgpack_parity() { + let body = serde_json::json!({ + "prompt": "Read", + "images": [{"data": "aGk"}], + "max_new_tokens": 8, + }); + assert!(generate_params_from_json(&body).is_err()); + let rmpv::Value::Map(map) = json_to_rmpv(body) else { + panic!("object must map") + }; + assert!(generate_params_from_rmpv(&map).is_err()); + } + + #[test] + fn test_generate_native_images_fail_closed_on_model_capability() { + let image_body = serde_json::json!({ + "prompt": "Read", + "images": [{"data": "aGk="}], + "max_new_tokens": 8, + }); + let image_params = _expect_generate_ok(&image_body); + assert!(!native_generate_image_input_allowed( + Some(&image_params), + false + )); + assert!(native_generate_image_input_allowed( + Some(&image_params), + true + )); + + let text_params = _expect_generate_ok(&serde_json::json!({ + "prompt": "Read", + "max_new_tokens": 8, + })); + assert!(native_generate_image_input_allowed( + Some(&text_params), + false + )); + } + + #[test] + fn test_generate_native_image_caps_fit_bounded_aggregate_body() { + assert!(native_generate_image_encoded_size_allowed( + MAX_GENERATE_IMAGE_BASE64_CHARS + )); + assert!(!native_generate_image_encoded_size_allowed(0)); + assert!(!native_generate_image_encoded_size_allowed( + MAX_GENERATE_IMAGE_BASE64_CHARS + 1 + )); + } + #[test] fn test_generate_params_from_json_rejects_empty_prompt() { let body = serde_json::json!({"prompt": "", "max_new_tokens": 8}); diff --git a/packages/sie_gateway/src/openapi.rs b/packages/sie_gateway/src/openapi.rs index 75bd66ef..bbac0023 100644 --- a/packages/sie_gateway/src/openapi.rs +++ b/packages/sie_gateway/src/openapi.rs @@ -122,6 +122,7 @@ static OPENAPI_JSON: LazyLock = LazyLock::new(|| { OpenAITranscriptionResponse, OpenAITranscriptionResponseFormat, OpenAITranscriptionUsage, + NativeGenerateImage, GenerateRequest, GenerateResponse, GenerateChunk, @@ -470,8 +471,9 @@ fn patch_generate_path(value: &mut Value) { post.insert( "description".to_string(), json!( - "SIE-native text generation. Omit `stream` or set it to false for a blocking JSON \ - response; set `stream: true` for SIE-native Server-Sent Events terminated by \ + "SIE-native generation with optional bounded inline images. Omit `stream` or set it \ + to false for a blocking JSON response; set `stream: true` for SIE-native \ + Server-Sent Events terminated by \ `data: [DONE]`. The model path parameter must use the SIE-safe ID (for example \ `Qwen__Qwen3-4B-Instruct`); HF-style slashes reject with 400." ), @@ -1750,11 +1752,29 @@ pub struct OpenAITranscriptionResponse { } // ── /v1/generate/{model} schemas ────────────────────────────────── +/// One inline image on the SIE-native generate surface. +#[derive(Debug, Serialize, Deserialize, ToSchema)] +pub struct NativeGenerateImage { + /// Canonical standard-base64 encoded image bytes, at most 16 MiB decoded. + /// Remote URLs are not accepted. + #[schema(min_length = 1, max_length = 22369624)] + pub data: String, + /// Short media-format hint such as ``png`` or ``jpeg``. + #[schema(min_length = 1, max_length = 32, pattern = "^[A-Za-z0-9.+-]+$")] + #[serde(default, skip_serializing_if = "Option::is_none")] + pub format: Option, +} + /// SIE-native text-generation request. Set ``stream`` to true for SSE. #[derive(Debug, Serialize, Deserialize, ToSchema)] pub struct GenerateRequest { #[schema(min_length = 1)] pub prompt: String, + /// Optional inline images paired with ``prompt``. When present, the worker + /// renders one user turn through the model's native chat template. + #[schema(min_items = 1, max_items = 16)] + #[serde(default, skip_serializing_if = "Option::is_none")] + pub images: Option>, #[schema(minimum = 1)] pub max_new_tokens: u32, #[serde(default, skip_serializing_if = "Option::is_none")] @@ -3134,6 +3154,19 @@ mod tests { let request = &spec["components"]["schemas"]["GenerateRequest"]; assert_eq!(request["required"], json!(["prompt", "max_new_tokens"])); + assert_eq!( + request["properties"]["images"]["items"]["$ref"], + "#/components/schemas/NativeGenerateImage" + ); + assert_eq!(request["properties"]["images"]["maxItems"], json!(16)); + assert_eq!( + spec["components"]["schemas"]["NativeGenerateImage"]["required"], + json!(["data"]) + ); + assert_eq!( + spec["components"]["schemas"]["NativeGenerateImage"]["properties"]["data"]["maxLength"], + json!(22_369_624) + ); let seed = &request["properties"]["seed"]; assert_eq!(seed["type"], json!(["integer", "null"])); assert_eq!(seed["format"], "int64"); diff --git a/packages/sie_sdk/src/sie_sdk/__init__.py b/packages/sie_sdk/src/sie_sdk/__init__.py index ef73ae55..8b7414dd 100644 --- a/packages/sie_sdk/src/sie_sdk/__init__.py +++ b/packages/sie_sdk/src/sie_sdk/__init__.py @@ -51,12 +51,15 @@ Connection, ConnectionCreated, ConnectionRevoked, + EbnfGrammar, EncodeResult, Entity, ExtractResult, File, FileDeleted, GenerateChunk, + GenerateGrammar, + GenerateImage, GenerateResult, GenerationUsage, HealthResponse, @@ -69,6 +72,7 @@ JobState, JobStatus, JobSubmitResult, + JsonSchemaGrammar, ModelInfo, ModelSummary, PoolInfo, @@ -77,6 +81,7 @@ PoolSpec, PoolSpecResponse, PoolStatusInfo, + RegexGrammar, ScoreResult, SparseResult, StatusMessage, @@ -102,12 +107,15 @@ "Connection", "ConnectionCreated", "ConnectionRevoked", + "EbnfGrammar", "EncodeResult", "Entity", "ExtractResult", "File", "FileDeleted", "GenerateChunk", + "GenerateGrammar", + "GenerateImage", "GenerateResult", "GenerationUsage", "HealthResponse", @@ -121,6 +129,7 @@ "JobState", "JobStatus", "JobSubmitResult", + "JsonSchemaGrammar", "LoraLoadingError", "ModelInfo", "ModelLoadFailedError", @@ -134,6 +143,7 @@ "PoolSpecResponse", "PoolStatusInfo", "ProvisioningError", + "RegexGrammar", "RequestError", "SIEAsyncClient", "SIEClient", diff --git a/packages/sie_sdk/src/sie_sdk/client/_shared.py b/packages/sie_sdk/src/sie_sdk/client/_shared.py index 9fb40201..83abe754 100644 --- a/packages/sie_sdk/src/sie_sdk/client/_shared.py +++ b/packages/sie_sdk/src/sie_sdk/client/_shared.py @@ -13,7 +13,7 @@ from email.utils import parsedate_to_datetime from importlib.metadata import PackageNotFoundError from importlib.metadata import version as pkg_version -from typing import Any, Protocol +from typing import Any, Protocol, cast from urllib.parse import urljoin, urlsplit import msgpack @@ -147,6 +147,7 @@ def json(self) -> Any: ... EntityResult, ExtractItemErrorDetail, ExtractResult, + GenerateGrammar, Relation, RequestMetadata, RequestUsage, @@ -179,6 +180,52 @@ def json(self) -> Any: ... DEFAULT_PROVISION_TIMEOUT_S = 900.0 # 15 minutes DEFAULT_RETRY_DELAY_S = 5.0 # Retry every 5 seconds if no Retry-After header +_GENERATE_GRAMMAR_VARIANTS = frozenset({"json_schema", "regex", "ebnf"}) +_GENERATE_GRAMMAR_FIELDS = _GENERATE_GRAMMAR_VARIANTS | {"label", "strict"} + + +def validate_generate_grammar(grammar: GenerateGrammar | Mapping[str, Any]) -> GenerateGrammar: + """Validate and detach the native structured-output grammar envelope. + + ``generate`` historically accepted a broad dictionary here. Keep that + source-compatible input type while validating the exact native three-arm + shape before issuing a request. + """ + if not isinstance(grammar, Mapping): + msg = "grammar must be a mapping" + raise TypeError(msg) + + unknown = set(grammar) - _GENERATE_GRAMMAR_FIELDS + if unknown: + names = ", ".join(sorted(str(name) for name in unknown)) + msg = f"grammar contains unsupported field(s): {names}" + raise ValueError(msg) + + variants = _GENERATE_GRAMMAR_VARIANTS.intersection(grammar) + if len(variants) != 1: + msg = "grammar must contain exactly one of json_schema, regex, or ebnf" + raise ValueError(msg) + + variant = next(iter(variants)) + value = grammar[variant] + if variant == "json_schema": + if not isinstance(value, Mapping): + msg = "grammar.json_schema must be a mapping" + raise TypeError(msg) + elif not isinstance(value, str): + msg = f"grammar.{variant} must be a string" + raise TypeError(msg) + + if "label" in grammar and not isinstance(grammar["label"], str): + msg = "grammar.label must be a string" + raise TypeError(msg) + if "strict" in grammar and not isinstance(grammar["strict"], bool): + msg = "grammar.strict must be a boolean" + raise TypeError(msg) + + return cast("GenerateGrammar", dict(grammar)) + + # Pool settings DEFAULT_LEASE_RENEWAL_INTERVAL_S = 60.0 # Renew lease every 60s (lease is 1200s) diff --git a/packages/sie_sdk/src/sie_sdk/client/async_.py b/packages/sie_sdk/src/sie_sdk/client/async_.py index c5f5b2e2..2a9b695b 100644 --- a/packages/sie_sdk/src/sie_sdk/client/async_.py +++ b/packages/sie_sdk/src/sie_sdk/client/async_.py @@ -46,7 +46,7 @@ from sie_sdk.audio import convert_item_audio from sie_sdk.documents import convert_item_document from sie_sdk.files import resolve_upload -from sie_sdk.images import convert_item_images +from sie_sdk.images import ImageLike, convert_images_for_json, convert_item_images from sie_sdk.jobs import TERMINAL_JOB_STATES, build_job_body, decode_chunk_bytes, job_chunks from sie_sdk.types import ( Batch, @@ -62,6 +62,8 @@ File, FileDeleted, GenerateChunk, + GenerateGrammar, + GenerateImage, GenerateResult, Item, JobResults, @@ -123,6 +125,7 @@ sse_chunk_error, sse_headers, validate_encode_result_count, + validate_generate_grammar, websocket_matches_base_url_origin, ) from ._sse import aiter_sse_payloads @@ -1848,12 +1851,13 @@ async def generate( prompt: str, *, max_new_tokens: int, + images: Sequence[ImageLike | GenerateImage] | None = None, temperature: float | None = None, top_p: float | None = None, stop: list[str] | None = None, frequency_penalty: float | None = None, presence_penalty: float | None = None, - grammar: dict[str, Any] | None = None, + grammar: GenerateGrammar | Mapping[str, Any] | None = None, seed: int | None = None, logit_bias: dict[str, float] | None = None, routing_key: str | None = None, @@ -1872,6 +1876,7 @@ async def generate( awaits the aggregated outcome; use :meth:`stream_generate` for SIE-native chunk streaming. """ + resolved_grammar = validate_generate_grammar(grammar) if grammar is not None else None pool_name, resolved_gpu = await self._resolve_pool_and_gpu(gpu) safe_model = model.replace("/", "__") @@ -1881,6 +1886,8 @@ async def generate( "prompt": prompt, "max_new_tokens": max_new_tokens, } + if images is not None: + request_body["images"] = convert_images_for_json(images) if stop is not None: request_body["stop"] = stop optional_fields = { @@ -1889,7 +1896,7 @@ async def generate( "options": resolved_options, "frequency_penalty": frequency_penalty, "presence_penalty": presence_penalty, - "grammar": grammar, + "grammar": resolved_grammar, "seed": seed, "logit_bias": logit_bias, "routing_key": routing_key, @@ -2266,12 +2273,13 @@ async def stream_generate( prompt: str, *, max_new_tokens: int, + images: Sequence[ImageLike | GenerateImage] | None = None, temperature: float | None = None, top_p: float | None = None, stop: list[str] | None = None, frequency_penalty: float | None = None, presence_penalty: float | None = None, - grammar: dict[str, Any] | None = None, + grammar: GenerateGrammar | Mapping[str, Any] | None = None, seed: int | None = None, logit_bias: dict[str, float] | None = None, logprobs: bool = False, @@ -2291,6 +2299,7 @@ async def stream_generate( Async counterpart of :meth:`SIEClient.stream_generate`. """ + resolved_grammar = validate_generate_grammar(grammar) if grammar is not None else None pool_name, resolved_gpu = await self._resolve_pool_and_gpu(gpu) safe_model = model.replace("/", "__") resolved_options = self._resolve_options(options) @@ -2299,6 +2308,8 @@ async def stream_generate( "max_new_tokens": max_new_tokens, "stream": True, } + if images is not None: + req["images"] = convert_images_for_json(images) if stop is not None: req["stop"] = stop optional_fields = { @@ -2307,7 +2318,7 @@ async def stream_generate( "top_p": top_p, "options": resolved_options, "presence_penalty": presence_penalty, - "grammar": grammar, + "grammar": resolved_grammar, "seed": seed, "logit_bias": logit_bias, "routing_key": routing_key, diff --git a/packages/sie_sdk/src/sie_sdk/client/sync.py b/packages/sie_sdk/src/sie_sdk/client/sync.py index 572e1328..57a391be 100644 --- a/packages/sie_sdk/src/sie_sdk/client/sync.py +++ b/packages/sie_sdk/src/sie_sdk/client/sync.py @@ -54,7 +54,7 @@ from sie_sdk.audio import convert_item_audio from sie_sdk.documents import convert_item_document from sie_sdk.files import resolve_upload -from sie_sdk.images import convert_item_images +from sie_sdk.images import ImageLike, convert_images_for_json, convert_item_images from sie_sdk.jobs import TERMINAL_JOB_STATES, build_job_body, decode_chunk_bytes, job_chunks from sie_sdk.types import ( Batch, @@ -70,6 +70,8 @@ File, FileDeleted, GenerateChunk, + GenerateGrammar, + GenerateImage, GenerateResult, Item, JobResults, @@ -132,6 +134,7 @@ sse_chunk_error, sse_headers, validate_encode_result_count, + validate_generate_grammar, websocket_matches_base_url_origin, ) from ._sse import iter_sse_payloads @@ -2030,12 +2033,13 @@ def generate( prompt: str, *, max_new_tokens: int, + images: Sequence[ImageLike | GenerateImage] | None = None, temperature: float | None = None, top_p: float | None = None, stop: list[str] | None = None, frequency_penalty: float | None = None, presence_penalty: float | None = None, - grammar: dict[str, Any] | None = None, + grammar: GenerateGrammar | Mapping[str, Any] | None = None, seed: int | None = None, logit_bias: dict[str, float] | None = None, routing_key: str | None = None, @@ -2065,6 +2069,9 @@ def generate( helpers in the SDK (use the OpenAI SDK against ``/v1/chat/completions`` for chat-shaped requests). max_new_tokens: Hard cap on output tokens. + images: Optional native image inputs. Pass encoded bytes, a path, + a PIL/NumPy image, or ``{"data": ..., "format": "png"}``. + The format hint is optional when the SDK can detect it. temperature: Sampling temperature override. Omit to use the selected model profile's default. top_p: Nucleus sampling cutoff override. Omit to use the selected @@ -2072,7 +2079,8 @@ def generate( stop: Optional list of stop strings. frequency_penalty: OpenAI-compatible frequency penalty in ``[-2, 2]``. presence_penalty: OpenAI-compatible presence penalty in ``[-2, 2]``. - grammar: Optional native structured-output grammar. + grammar: Optional native structured-output grammar. Set exactly + one of ``json_schema``, ``regex``, or ``ebnf``. seed: Optional signed 64-bit per-request sampling seed. Exact reproducibility depends on the active generation backend and deployment configuration. @@ -2104,6 +2112,7 @@ def generate( retried) or other 5xx responses. """ self._reset_retry_count() + resolved_grammar = validate_generate_grammar(grammar) if grammar is not None else None pool_name, resolved_gpu = self._resolve_pool_and_gpu(gpu) safe_model = model.replace("/", "__") @@ -2113,6 +2122,8 @@ def generate( "prompt": prompt, "max_new_tokens": max_new_tokens, } + if images is not None: + request_body["images"] = convert_images_for_json(images) if stop is not None: request_body["stop"] = stop optional_fields = { @@ -2121,7 +2132,7 @@ def generate( "options": resolved_options, "frequency_penalty": frequency_penalty, "presence_penalty": presence_penalty, - "grammar": grammar, + "grammar": resolved_grammar, "seed": seed, "logit_bias": logit_bias, "routing_key": routing_key, @@ -2494,12 +2505,13 @@ def stream_generate( prompt: str, *, max_new_tokens: int, + images: Sequence[ImageLike | GenerateImage] | None = None, temperature: float | None = None, top_p: float | None = None, stop: list[str] | None = None, frequency_penalty: float | None = None, presence_penalty: float | None = None, - grammar: dict[str, Any] | None = None, + grammar: GenerateGrammar | Mapping[str, Any] | None = None, seed: int | None = None, logit_bias: dict[str, float] | None = None, logprobs: bool = False, @@ -2521,6 +2533,7 @@ def stream_generate( ``done: true`` plus ``usage`` / ``ttft_ms``. Error semantics match :meth:`stream_chat_completions`. """ + resolved_grammar = validate_generate_grammar(grammar) if grammar is not None else None pool_name, resolved_gpu = self._resolve_pool_and_gpu(gpu) safe_model = model.replace("/", "__") resolved_options = self._resolve_options(options) @@ -2529,6 +2542,8 @@ def stream_generate( "max_new_tokens": max_new_tokens, "stream": True, } + if images is not None: + req["images"] = convert_images_for_json(images) if stop is not None: req["stop"] = stop optional_fields = { @@ -2537,7 +2552,7 @@ def stream_generate( "top_p": top_p, "options": resolved_options, "presence_penalty": presence_penalty, - "grammar": grammar, + "grammar": resolved_grammar, "seed": seed, "logit_bias": logit_bias, "routing_key": routing_key, diff --git a/packages/sie_sdk/src/sie_sdk/images.py b/packages/sie_sdk/src/sie_sdk/images.py index 71d02eb6..f396e30d 100644 --- a/packages/sie_sdk/src/sie_sdk/images.py +++ b/packages/sie_sdk/src/sie_sdk/images.py @@ -1,14 +1,19 @@ """Image conversion utilities for SIE SDK. -Images are serialized as JPEG bytes for transport. -This module handles conversion from various input formats to JPEG bytes. +PIL/array/path inputs are serialized as JPEG bytes for transport. Already +encoded bytes remain byte-identical and carry their detected format. -Wire format: raw JPEG bytes in msgpack (no base64 encoding). +MessagePack primitives carry raw image bytes plus a truthful format token. +Already encoded JPEG, PNG, GIF, WebP, BMP, and TIFF bytes remain native; +unknown signatures (including currently unsupported HEIC/AVIF) fail closed. +Native JSON generation base64-encodes the same bytes at its final wire step. """ from __future__ import annotations +import base64 import io +from collections.abc import Sequence from pathlib import Path from typing import TYPE_CHECKING, Any, Union @@ -22,6 +27,60 @@ # Default JPEG quality for image transport. DEFAULT_JPEG_QUALITY = 95 +_MAX_IMAGE_FORMAT_LENGTH = 32 +_MAX_GENERATE_IMAGES = 16 +_MAX_GENERATE_IMAGE_BYTES = 16 * 1024 * 1024 + + +def _canonical_image_format(value: object) -> str: + if not isinstance(value, str) or not ( + 1 <= len(value) <= _MAX_IMAGE_FORMAT_LENGTH + and all(character.isascii() and (character.isalnum() or character in ".+-") for character in value) + ): + msg = "Image format must be a short ASCII media-format token" + raise ValueError(msg) + normalized = value.lower() + return "jpeg" if normalized in {"jpg", "jpe"} else normalized + + +def _detect_encoded_image_format(data: bytes) -> str: + """Detect the transport format of already-encoded image bytes. + + This deliberately inspects signatures rather than decoding pixels: SDK + serialization must stay cheap, while the server-side media loader remains + responsible for full image validation. Unknown bytes fail closed instead + of being silently labeled as JPEG. + """ + if data.startswith(b"\xff\xd8\xff"): + return "jpeg" + if data.startswith(b"\x89PNG\r\n\x1a\n"): + return "png" + if data.startswith((b"GIF87a", b"GIF89a")): + return "gif" + if len(data) >= 12 and data.startswith(b"RIFF") and data[8:12] == b"WEBP": + return "webp" + if data.startswith(b"BM"): + return "bmp" + if data.startswith((b"II*\x00", b"MM\x00*", b"II+\x00", b"MM\x00+")): + return "tiff" + msg = "Could not detect encoded image format from bytes" + raise ValueError(msg) + + +def _image_wire_value(image: ImageLike, declared_format: object = None) -> dict[str, Any]: + if isinstance(image, bytes): + data = image + detected_format = _detect_encoded_image_format(data) + else: + data = to_jpeg_bytes(image) + detected_format = "jpeg" + + if declared_format is not None: + normalized_format = _canonical_image_format(declared_format) + if normalized_format != detected_format: + msg = f"Image format mismatch: declared '{normalized_format}', detected '{detected_format}'" + raise ValueError(msg) + return {"data": data, "format": detected_format} def to_jpeg_bytes( @@ -117,8 +176,9 @@ def _pil_to_jpeg_bytes(image: Image.Image, *, quality: int = DEFAULT_JPEG_QUALIT def convert_item_images(item: dict[str, Any]) -> dict[str, Any]: """Convert all images in an item to wire format for transport. - Images are sent as JPEG bytes wrapped in - ImageInput format: {"data": , "format": "jpeg"}. + Images are wrapped in ImageInput format: ``{"data": , + "format": }``. PIL/array/path inputs are + converted to JPEG; encoded bytes are preserved and signature-detected. Modifies the item in-place and returns it. @@ -139,13 +199,34 @@ def convert_item_images(item: dict[str, Any]) -> dict[str, Any]: for img in images: # Handle ImageInput dict format (SDK user provided dict with "data" key) if isinstance(img, dict) and "data" in img: - img_data = img["data"] - jpeg_bytes = to_jpeg_bytes(img_data) - converted.append({"data": jpeg_bytes, "format": "jpeg"}) + converted.append(_image_wire_value(img["data"], img.get("format"))) else: # Direct image input (PIL.Image, ndarray, bytes, str/Path) - jpeg_bytes = to_jpeg_bytes(img) - converted.append({"data": jpeg_bytes, "format": "jpeg"}) + converted.append(_image_wire_value(img)) item["images"] = converted return item + + +def convert_images_for_json(images: Sequence[Any]) -> list[dict[str, str]]: + """Convert native image inputs to the JSON generate wire envelope. + + Native tensor primitives carry image bytes directly in MessagePack. The + native generate endpoint is JSON, so it uses the same ``{data, format}`` + envelope with standard-base64 text in ``data``. Conversion is centralized + here so sync, async, buffered, and streaming clients cannot drift. + """ + if not 1 <= len(images) <= _MAX_GENERATE_IMAGES: + msg = f"images must contain between 1 and {_MAX_GENERATE_IMAGES} entries" + raise ValueError(msg) + converted = convert_item_images({"images": list(images)}).get("images", []) + if any(not 1 <= len(image["data"]) <= _MAX_GENERATE_IMAGE_BYTES for image in converted): + msg = "each image must contain between 1 byte and 16 MiB" + raise ValueError(msg) + return [ + { + "data": base64.b64encode(image["data"]).decode("ascii"), + "format": str(image.get("format") or "jpeg"), + } + for image in converted + ] diff --git a/packages/sie_sdk/src/sie_sdk/types.py b/packages/sie_sdk/src/sie_sdk/types.py index 77e9ff64..256e374e 100644 --- a/packages/sie_sdk/src/sie_sdk/types.py +++ b/packages/sie_sdk/src/sie_sdk/types.py @@ -10,7 +10,7 @@ from __future__ import annotations from collections.abc import Sequence -from typing import TYPE_CHECKING, Any, Literal, NotRequired, Required, TypedDict +from typing import TYPE_CHECKING, Any, Literal, Never, NotRequired, Required, TypedDict import numpy as np @@ -386,6 +386,47 @@ class ExtractResult(TypedDict, total=False): FinishReason = Literal["stop", "length", "cancelled", "content_filter", "error"] +class GenerateImage(TypedDict): + """One image paired with a native generation prompt.""" + + data: Image.Image | NDArray[Any] | bytes | str | Path + format: NotRequired[str] + + +class JsonSchemaGrammar(TypedDict): + """Constrain generation to a JSON Schema.""" + + json_schema: dict[str, Any] + regex: NotRequired[Never] + ebnf: NotRequired[Never] + label: NotRequired[str] + strict: NotRequired[bool] + + +class RegexGrammar(TypedDict): + """Constrain generation to a regular expression.""" + + json_schema: NotRequired[Never] + regex: str + ebnf: NotRequired[Never] + label: NotRequired[str] + strict: NotRequired[bool] + + +class EbnfGrammar(TypedDict): + """Constrain generation to an EBNF grammar.""" + + json_schema: NotRequired[Never] + regex: NotRequired[Never] + ebnf: str + label: NotRequired[str] + strict: NotRequired[bool] + + +GenerateGrammar = JsonSchemaGrammar | RegexGrammar | EbnfGrammar +"""Native structured-output grammar. Exactly one grammar variant is set.""" + + class GenerationUsage(TypedDict): """Token usage for a single generation call.""" diff --git a/packages/sie_sdk/tests/client/test_generate.py b/packages/sie_sdk/tests/client/test_generate.py index 0899471a..52a010cd 100644 --- a/packages/sie_sdk/tests/client/test_generate.py +++ b/packages/sie_sdk/tests/client/test_generate.py @@ -11,6 +11,7 @@ from __future__ import annotations import json +from typing import Any from unittest.mock import AsyncMock, MagicMock, patch import httpx @@ -130,6 +131,62 @@ async def _aenter() -> MagicMock: class TestSyncGenerate: + def test_generate_serializes_images_and_typed_json_schema_grammar(self) -> None: + grammar = { + "json_schema": {"type": "object", "properties": {"title": {"type": "string"}}}, + "label": "document", + "strict": True, + } + with patch("sie_sdk.client.sync.httpx.Client") as mock_client: + mock_client.return_value.post.return_value = _ok_response(_ok_envelope()) + client = SIEClient("http://localhost:8080") + client.generate( + "vision-model", + prompt="Read this document", + images=[{"data": b"\x89PNG\r\n\x1a\npayload"}], + grammar=grammar, + max_new_tokens=32, + ) + + sent = json.loads(mock_client.return_value.post.call_args.kwargs["content"]) + assert sent["images"] == [{"data": "iVBORw0KGgpwYXlsb2Fk", "format": "png"}] + assert sent["grammar"] == grammar + client.close() + + @pytest.mark.parametrize( + "grammar", + [ + {"json_schema": {"type": "object"}}, + {"regex": r"\d{3}-\d{4}"}, + {"ebnf": 'root ::= "yes" | "no"'}, + ], + ) + def test_generate_preserves_each_typed_grammar_variant(self, grammar: dict[str, Any]) -> None: + with patch("sie_sdk.client.sync.httpx.Client") as mock_client: + mock_client.return_value.post.return_value = _ok_response(_ok_envelope()) + client = SIEClient("http://localhost:8080") + client.generate("m", prompt="Hi", grammar=grammar, max_new_tokens=8) + sent = json.loads(mock_client.return_value.post.call_args.kwargs["content"]) + assert sent["grammar"] == grammar + client.close() + + @pytest.mark.parametrize( + "grammar", + [ + {}, + {"json_schema": {}, "regex": "x"}, + {"regex": 123}, + {"ebnf": "root", "unknown": True}, + ], + ) + def test_generate_rejects_invalid_grammar_before_request(self, grammar: dict[str, Any]) -> None: + with patch("sie_sdk.client.sync.httpx.Client") as mock_client: + client = SIEClient("http://localhost:8080") + with pytest.raises((TypeError, ValueError)): + client.generate("m", prompt="Hi", grammar=grammar, max_new_tokens=8) + mock_client.return_value.post.assert_not_called() + client.close() + def test_generate_happy_path_parses_envelope(self) -> None: envelope = { "model": "Qwen__Qwen3-4B-Instruct-2507", @@ -171,6 +228,9 @@ def test_generate_happy_path_parses_envelope(self) -> None: sent = json.loads(mock_client.return_value.post.call_args.kwargs["content"].decode("utf-8")) assert "temperature" not in sent assert "top_p" not in sent + assert mock_client.return_value.post.call_args.kwargs["content"] == ( + b'{"prompt": "Hi", "max_new_tokens": 32}' + ) client.close() @@ -387,6 +447,36 @@ def test_generate_does_not_retry_mid_flight_transport_error(self) -> None: class TestAsyncGenerate: + @pytest.mark.asyncio + async def test_generate_serializes_native_image_and_ebnf_grammar(self) -> None: + raw = MagicMock() + raw.status = 200 + raw.headers = {} + raw.read = AsyncMock(return_value=json.dumps(_ok_envelope()).encode()) + post_ctx = MagicMock() + post_ctx.__aenter__ = AsyncMock(return_value=raw) + post_ctx.__aexit__ = AsyncMock(return_value=None) + session = MagicMock() + session.post = MagicMock(return_value=post_ctx) + session.close = AsyncMock() + + with patch("sie_sdk.client.async_.aiohttp.ClientSession", return_value=session): + client = SIEAsyncClient("http://localhost:8080") + try: + await client.generate( + "vision-model", + prompt="Answer yes or no", + images=[b"\xff\xd8\xff\xe0hello"], + grammar={"ebnf": 'root ::= "yes" | "no"'}, + max_new_tokens=8, + ) + finally: + await client.close() + + sent = json.loads(session.post.call_args.kwargs["data"]) + assert sent["images"] == [{"data": "/9j/4GhlbGxv", "format": "jpeg"}] + assert sent["grammar"] == {"ebnf": 'root ::= "yes" | "no"'} + @pytest.mark.asyncio async def test_generate_non_dict_response_retains_request_metadata(self) -> None: client = SIEAsyncClient("http://localhost:8080") diff --git a/packages/sie_sdk/tests/test_images.py b/packages/sie_sdk/tests/test_images.py index b9c3b287..68ee00ac 100644 --- a/packages/sie_sdk/tests/test_images.py +++ b/packages/sie_sdk/tests/test_images.py @@ -7,7 +7,7 @@ import numpy as np import pytest from PIL import Image -from sie_sdk.images import convert_item_images, to_jpeg_bytes +from sie_sdk.images import convert_images_for_json, convert_item_images, to_jpeg_bytes class TestToJpegBytes: @@ -182,3 +182,34 @@ def test_mixed_formats(self) -> None: for img_input in result["images"]: assert isinstance(img_input, dict) assert isinstance(img_input["data"], bytes) + + def test_raw_png_bytes_preserve_detected_format(self) -> None: + png_bytes = b"\x89PNG\r\n\x1a\npayload" + + result = convert_item_images({"images": [png_bytes, {"data": png_bytes, "format": "PNG"}]}) + + assert result["images"] == [ + {"data": png_bytes, "format": "png"}, + {"data": png_bytes, "format": "png"}, + ] + + @pytest.mark.parametrize("declared_format", ["jpeg", "png;url=https://example.com", 7]) + def test_declared_format_mismatch_or_invalid_token_fails_closed(self, declared_format: object) -> None: + png_bytes = b"\x89PNG\r\n\x1a\npayload" + + with pytest.raises(ValueError, match="Image format"): + convert_item_images({"images": [{"data": png_bytes, "format": declared_format}]}) + + def test_unknown_raw_bytes_fail_closed(self) -> None: + with pytest.raises(ValueError, match="Could not detect encoded image format"): + convert_item_images({"images": [b"not-an-image"]}) + + +class TestConvertImagesForJson: + def test_base64_encodes_detected_image(self) -> None: + assert convert_images_for_json([b"\xff\xd8\xff\xe0hello"]) == [{"data": "/9j/4GhlbGxv", "format": "jpeg"}] + + @pytest.mark.parametrize("images", [[], [b"\xff\xd8\xff"] * 17]) + def test_rejects_invalid_image_count(self, images: list[bytes]) -> None: + with pytest.raises(ValueError, match="between 1 and 16"): + convert_images_for_json(images) diff --git a/packages/sie_server/openapi.json b/packages/sie_server/openapi.json index d46352d5..2952461f 100644 --- a/packages/sie_server/openapi.json +++ b/packages/sie_server/openapi.json @@ -2446,6 +2446,157 @@ "title": "GenerateChunk", "type": "object" }, + "NativeEbnfGrammarModel": { + "additionalProperties": false, + "description": "EBNF structured-output constraint.", + "properties": { + "ebnf": { + "title": "Ebnf", + "type": "string" + }, + "label": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Label" + }, + "strict": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Strict" + } + }, + "required": [ + "ebnf" + ], + "title": "NativeEbnfGrammarModel", + "type": "object" + }, + "NativeGenerateImageModel": { + "description": "One inline image on the SIE-native generate surface.", + "properties": { + "data": { + "description": "Canonical standard-base64 image bytes, at most 16 MiB decoded", + "maxLength": 22369624, + "minLength": 1, + "title": "Data", + "type": "string" + }, + "format": { + "anyOf": [ + { + "maxLength": 32, + "minLength": 1, + "pattern": "^[A-Za-z0-9.+-]+$", + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "description": "Short media-format hint such as png or jpeg", + "title": "Format" + } + }, + "required": [ + "data" + ], + "title": "NativeGenerateImageModel", + "type": "object" + }, + "NativeJsonSchemaGrammarModel": { + "additionalProperties": false, + "description": "JSON Schema structured-output constraint.", + "properties": { + "json_schema": { + "additionalProperties": true, + "title": "Json Schema", + "type": "object" + }, + "label": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Label" + }, + "strict": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Strict" + } + }, + "required": [ + "json_schema" + ], + "title": "NativeJsonSchemaGrammarModel", + "type": "object" + }, + "NativeRegexGrammarModel": { + "additionalProperties": false, + "description": "Regular-expression structured-output constraint.", + "properties": { + "regex": { + "title": "Regex", + "type": "string" + }, + "label": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Label" + }, + "strict": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Strict" + } + }, + "required": [ + "regex" + ], + "title": "NativeRegexGrammarModel", + "type": "object" + }, "GenerateRequestModel": { "description": "Request body for the SIE-native generate endpoint.", "examples": [ @@ -2464,6 +2615,43 @@ "title": "Prompt", "type": "string" }, + "images": { + "anyOf": [ + { + "items": { + "$ref": "#/components/schemas/NativeGenerateImageModel" + }, + "maxItems": 16, + "minItems": 1, + "type": "array" + }, + { + "type": "null" + } + ], + "default": null, + "description": "Optional inline images paired with prompt. The worker renders one user turn through the model's native chat template; remote URLs are not accepted.", + "title": "Images" + }, + "grammar": { + "anyOf": [ + { + "$ref": "#/components/schemas/NativeJsonSchemaGrammarModel" + }, + { + "$ref": "#/components/schemas/NativeRegexGrammarModel" + }, + { + "$ref": "#/components/schemas/NativeEbnfGrammarModel" + }, + { + "type": "null" + } + ], + "default": null, + "description": "Optional structured-output grammar; exactly one of json_schema, regex, or ebnf", + "title": "Grammar" + }, "max_new_tokens": { "description": "Maximum number of tokens to generate", "minimum": 1, diff --git a/packages/sie_server/src/sie_server/adapters/_generation_base.py b/packages/sie_server/src/sie_server/adapters/_generation_base.py index 16d7dfe9..6fdfed08 100644 --- a/packages/sie_server/src/sie_server/adapters/_generation_base.py +++ b/packages/sie_server/src/sie_server/adapters/_generation_base.py @@ -23,6 +23,7 @@ from sie_server.adapters._spec import AdapterSpec from sie_server.adapters.base import ModelAdapter, ModelCapabilities, ModelDims +from sie_server.types.grammar import GrammarSpec from sie_server.types.inputs import ImageInput logger = logging.getLogger(__name__) @@ -245,6 +246,7 @@ def generate( top_k: int | None = None, repetition_penalty: float | None = None, min_new_tokens: int | None = None, + grammar: GrammarSpec | None = None, seed: int | None = None, logit_bias: dict[str, float] | None = None, logprobs: bool = False, @@ -270,6 +272,7 @@ def generate( default (typically 0.0). Gateway-validated upstream. min_new_tokens: Optional minimum generated-token floor. Adapters that cannot enforce it must reject rather than silently ignore it. + grammar: Optional structured-output grammar. presence_penalty: Optional OpenAI-style presence penalty in ``[-2.0, 2.0]``. Same semantics as ``frequency_penalty``. diff --git a/packages/sie_server/src/sie_server/adapters/fake/adapter.py b/packages/sie_server/src/sie_server/adapters/fake/adapter.py index df129d17..857cac33 100644 --- a/packages/sie_server/src/sie_server/adapters/fake/adapter.py +++ b/packages/sie_server/src/sie_server/adapters/fake/adapter.py @@ -48,6 +48,7 @@ if TYPE_CHECKING: from collections.abc import AsyncIterator + from sie_server.types.grammar import GrammarSpec from sie_server.types.inputs import ImageInput, Item # Default declared footprint. Consumed by the MemoryManager once the @@ -409,6 +410,7 @@ async def generate( top_k: int | None = None, repetition_penalty: float | None = None, min_new_tokens: int | None = None, + grammar: GrammarSpec | None = None, seed: int | None = None, logit_bias: dict[str, float] | None = None, logprobs: bool = False, @@ -423,6 +425,7 @@ async def generate( presence_penalty, top_k, repetition_penalty, + grammar, seed, logit_bias, min_new_tokens, diff --git a/packages/sie_server/src/sie_server/api/generate.py b/packages/sie_server/src/sie_server/api/generate.py index f6243835..b275e41f 100644 --- a/packages/sie_server/src/sie_server/api/generate.py +++ b/packages/sie_server/src/sie_server/api/generate.py @@ -1,4 +1,4 @@ -"""Direct HTTP route for blocking text generation (walking-skeleton local-dev path). +"""Direct HTTP route for native generation (walking-skeleton local-dev path). This is the **local-dev** counterpart of the gateway's ``proxy_generate`` — it bypasses NATS/JetStream entirely and calls the @@ -23,8 +23,9 @@ .. code-block:: json - { "prompt": "...", "max_new_tokens": 64, "temperature": 0.7, - "top_p": 0.9, "stop": [""] } + { "prompt": "...", "images": [{"data": "", "format": "png"}], + "max_new_tokens": 64, "temperature": 0.7, "top_p": 0.9, + "stop": [""] } Response shape:: @@ -38,14 +39,20 @@ from __future__ import annotations +import asyncio +import base64 +import binascii import json import logging import math import os +import threading import time import uuid from collections.abc import AsyncIterator -from typing import Annotated, Any +from functools import lru_cache +from pathlib import Path +from typing import Annotated, Any, cast from fastapi import APIRouter, Header, HTTPException, Request, status from fastapi.responses import JSONResponse, StreamingResponse @@ -55,7 +62,10 @@ from sie_server.api.helpers import ModelStateChecker from sie_server.api.validation import validate_machine_profile_header, validate_signed_i64 from sie_server.core.runtime_options import apply_generation_runtime_options +from sie_server.core.tokenizer import image_first_chat_message, load_tokenizer from sie_server.observability.tracing import tracer +from sie_server.types.grammar import GrammarSpec +from sie_server.types.inputs import ImageInput from sie_server.types.openapi import ( GenerateInputTooLongErrorResponse, GenerateModelLoadFailedErrorResponse, @@ -90,6 +100,7 @@ # ``prompt_cache_key`` / ``safety_identifier``. _SUPPORTED_FIELDS = { "prompt", + "images", "max_new_tokens", "temperature", "top_p", @@ -97,6 +108,7 @@ "stream", "frequency_penalty", "presence_penalty", + "grammar", "seed", "logit_bias", "logprobs", @@ -108,16 +120,34 @@ } -# Maximum prompt size accepted by this direct route, in UTF-8 bytes. -# Mirrors the gateway's per-endpoint generate body cap -# (``MAX_GENERATE_BODY = 4 MiB`` in ``proxy.rs``): generate is pure text, -# Qwen3.5's 32k context is ~128 KiB of UTF-8, so 4 MiB is ~30× headroom -# while closing the trivial-OOM-under-concurrency vector. The gateway caps -# the whole body; this worker-local dev route never sits behind the -# gateway, so without this cap an oversized prompt would be deserialised, -# tokenised, and forwarded unbounded. Override via -# ``SIE_GENERATE_MAX_PROMPT_BYTES``. +# Maximum prompt size accepted by this direct route, in UTF-8 bytes. The +# gateway caps the whole body at 24 MiB so one 16 MiB decoded inline image +# fits after base64 expansion. This worker-local dev route never sits behind +# the gateway, so it independently caps both the prompt and request body. +# Override the prompt limit via ``SIE_GENERATE_MAX_PROMPT_BYTES``. _MAX_PROMPT_BYTES = int(os.environ.get("SIE_GENERATE_MAX_PROMPT_BYTES", str(4 * 1024 * 1024))) +_MAX_GENERATE_BODY_BYTES = int(os.environ.get("SIE_GENERATE_MAX_BODY_BYTES", str(24 * 1024 * 1024))) +_MAX_GENERATE_IMAGES = 16 +_MAX_GENERATE_IMAGE_BYTES = 16 * 1024 * 1024 +_MAX_GENERATE_IMAGE_BASE64_CHARS = 4 * ((_MAX_GENERATE_IMAGE_BYTES + 2) // 3) +_MAX_GENERATE_IMAGE_FORMAT_LENGTH = 32 +_MAX_GRAMMAR_BYTES = 64 * 1024 +_MAX_SCHEMA_DEPTH = 16 +_MAX_SCHEMA_NODES = 16 * 1024 +_MAX_REGEX_LENGTH = 4 * 1024 +_MAX_EBNF_LENGTH = 8 * 1024 +_UNSUPPORTED_SCHEMA_KEYWORDS = frozenset( + { + "$dynamicRef", + "if", + "then", + "else", + "unevaluatedProperties", + "dependentSchemas", + } +) +_NATIVE_TOKENIZER_CACHE_SIZE = 16 +_NATIVE_TOKENIZER_LOAD_LOCK = threading.Lock() # OpenAI penalty range (mirrors the gateway's ``proxy.rs::parse_penalty``): # ``frequency_penalty`` / ``presence_penalty`` must be a finite number in @@ -243,6 +273,305 @@ def _validate_logprobs(logprobs_value: Any, top_logprobs_value: Any) -> tuple[bo return bool(logprobs_enabled), top_logprobs_value +def _parse_native_images(value: Any) -> list[ImageInput] | None: + """Validate and decode the native JSON image envelope.""" + if value is None: + return None + if not isinstance(value, list) or not (1 <= len(value) <= _MAX_GENERATE_IMAGES): + raise _bad_request( + f"'images' must contain between 1 and {_MAX_GENERATE_IMAGES} entries", + param="images", + ) + images: list[ImageInput] = [] + for index, entry in enumerate(value): + owner = f"images[{index}]" + if not isinstance(entry, dict): + raise _bad_request(f"'{owner}' must be an object", param=owner) + entry_dict = cast("dict[str, Any]", entry) + unknown = set(entry_dict) - {"data", "format"} + if unknown: + param = f"{owner}.{sorted(unknown)[0]}" + raise _bad_request(f"'{param}' is not supported", param=param, code="unsupported_field") + encoded = entry_dict.get("data") + data_owner = f"{owner}.data" + if not isinstance(encoded, str) or not encoded: + raise _bad_request(f"'{data_owner}' must be a non-empty base64 string", param=data_owner) + if len(encoded) > _MAX_GENERATE_IMAGE_BASE64_CHARS: + raise _bad_request(f"'{data_owner}' exceeds the 16 MiB decoded-image limit", param=data_owner) + try: + data = base64.b64decode(encoded, validate=True) + except (binascii.Error, ValueError) as exc: + raise _bad_request(f"'{data_owner}' must be valid standard base64", param=data_owner) from exc + if base64.b64encode(data).decode("ascii") != encoded: + raise _bad_request(f"'{data_owner}' must use canonical standard base64", param=data_owner) + if not data or len(data) > _MAX_GENERATE_IMAGE_BYTES: + raise _bad_request( + f"'{data_owner}' must decode to between 1 byte and 16 MiB", + param=data_owner, + ) + format_value = entry_dict.get("format") + if format_value is not None and ( + not isinstance(format_value, str) + or not (1 <= len(format_value) <= _MAX_GENERATE_IMAGE_FORMAT_LENGTH) + or not all( + character.isascii() and (character.isalnum() or character in ".+-") for character in format_value + ) + ): + format_owner = f"{owner}.format" + raise _bad_request(f"'{format_owner}' must be a short media-format token", param=format_owner) + images.append({"data": data, "format": format_value.lower() if format_value else None}) + return images + + +def _schema_child_context(parent: str, key: str) -> str: + if parent == "schema": + if key in {"properties", "patternProperties", "$defs", "definitions", "dependentSchemas"}: + return "schema_map" + if key in {"oneOf", "anyOf", "allOf", "prefixItems"}: + return "schema_array" + if key in { + "items", + "additionalProperties", + "contains", + "propertyNames", + "not", + "if", + "then", + "else", + }: + return "schema" + return "other" + if parent == "schema_map": + return "schema" + return "other" + + +def _json_pointer(root: Any, pointer: str) -> Any: + current = root + if not pointer: + return current + for raw_token in pointer.removeprefix("/").split("/"): + token = raw_token.replace("~1", "/").replace("~0", "~") + if isinstance(current, dict) and token in current: + current = current[token] + elif isinstance(current, list) and token.isdigit() and int(token) < len(current): + current = current[int(token)] + else: + raise KeyError(pointer) + return current + + +def _dereference_schema_refs(schema: dict[str, Any]) -> dict[str, Any]: + visited = 0 + stack: list[str] = [] + + def resolve(value: Any, path: str, context: str) -> Any: + nonlocal visited + visited += 1 + if visited > _MAX_SCHEMA_NODES: + raise _bad_request( + f"JSON Schema node count exceeds limit ({_MAX_SCHEMA_NODES})", + param=path, + ) + if isinstance(value, dict): + if context == "schema" and "$ref" in value: + ref = value["$ref"] + ref_path = f"{path}.$ref" + if not isinstance(ref, str): + raise _bad_request("'$ref' must be a string", param=ref_path) + if not ref.startswith("#"): + raise _bad_request("external '$ref' is not supported", param=ref_path, code="unsupported_field") + pointer = ref[1:] + if pointer and not pointer.startswith("/"): + raise _bad_request( + "only internal JSON-pointer '$ref' values are supported", + param=ref_path, + code="unsupported_field", + ) + if pointer in stack: + raise _bad_request(f"recursive '$ref' cycle detected at {ref!r}", param=ref_path) + try: + target = _json_pointer(schema, pointer) + except KeyError as exc: + raise _bad_request(f"unresolved internal '$ref' {ref!r}", param=ref_path) from exc + stack.append(pointer) + resolved = resolve(target, path, "schema") + stack.pop() + siblings = { + key: resolve(child, f"{path}.{key}", _schema_child_context(context, key)) + for key, child in value.items() + if key not in {"$ref", "$defs", "definitions"} + } + return resolved if not siblings else {"allOf": [resolved, siblings]} + + return { + key: resolve(child, f"{path}.{key}", _schema_child_context(context, key)) + for key, child in value.items() + if not (context == "schema" and key in {"$defs", "definitions"}) + } + if isinstance(value, list): + child_context = "schema" if context in {"schema", "schema_array"} else "other" + return [resolve(child, f"{path}[{index}]", child_context) for index, child in enumerate(value)] + return value + + resolved = resolve(schema, "grammar.json_schema", "schema") + return cast("dict[str, Any]", resolved) + + +def _validate_schema_shape(schema: Any) -> None: + visited = 0 + nesting_keys = { + "properties", + "patternProperties", + "additionalProperties", + "unevaluatedProperties", + "items", + "prefixItems", + "contains", + "propertyNames", + "oneOf", + "anyOf", + "allOf", + "not", + "definitions", + "$defs", + "dependentSchemas", + "if", + "then", + "else", + } + + def walk(value: Any, path: str, depth: int) -> None: + nonlocal visited + visited += 1 + if visited > _MAX_SCHEMA_NODES: + raise _bad_request( + f"JSON Schema node count exceeds limit ({_MAX_SCHEMA_NODES})", + param=path, + ) + if depth > _MAX_SCHEMA_DEPTH: + raise _bad_request(f"JSON Schema depth exceeds limit ({_MAX_SCHEMA_DEPTH})", param=path) + if isinstance(value, dict): + unsupported = _UNSUPPORTED_SCHEMA_KEYWORDS.intersection(value) + if unsupported: + keyword = sorted(unsupported)[0] + raise _bad_request( + f"JSON Schema keyword '{keyword}' is not supported", + param=f"{path}.{keyword}", + code="unsupported_field", + ) + for key, child in value.items(): + walk(child, f"{path}.{key}", depth + 1 if key in nesting_keys else depth) + elif isinstance(value, list): + for index, child in enumerate(value): + walk(child, f"{path}[{index}]", depth) + + walk(schema, "grammar.json_schema", 0) + + +def _parse_native_grammar(value: Any) -> GrammarSpec | None: + """Validate the public native grammar envelope and build the adapter type.""" + if value is None: + return None + if not isinstance(value, dict): + raise _bad_request("'grammar' must be a JSON object", param="grammar") + encoded_size = len(json.dumps(value, separators=(",", ":")).encode("utf-8")) + if encoded_size > _MAX_GRAMMAR_BYTES: + raise _bad_request( + f"grammar payload {encoded_size} bytes exceeds limit ({_MAX_GRAMMAR_BYTES} bytes)", + param="grammar", + ) + unknown = set(value) - {"json_schema", "regex", "ebnf", "label", "strict"} + if unknown: + field = sorted(unknown)[0] + raise _bad_request(f"'grammar.{field}' is not supported", param=f"grammar.{field}", code="unsupported_field") + kinds = [kind for kind in ("json_schema", "regex", "ebnf") if kind in value] + if len(kinds) != 1: + raise _bad_request( + "'grammar' must contain exactly one of 'json_schema', 'regex' or 'ebnf'", + param="grammar", + ) + label = value.get("label") + if label is not None and not isinstance(label, str): + raise _bad_request("'grammar.label' must be a string", param="grammar.label") + strict = value.get("strict") + if strict is not None and not isinstance(strict, bool): + raise _bad_request("'grammar.strict' must be a boolean", param="grammar.strict") + + kind = kinds[0] + payload = value[kind] + if kind == "json_schema": + if not isinstance(payload, dict): + raise _bad_request("'grammar.json_schema' must be an object", param="grammar.json_schema") + payload = _dereference_schema_refs(cast("dict[str, Any]", payload)) + _validate_schema_shape(payload) + resolved_size = len( + json.dumps( + { + "json_schema": payload, + **({"label": label} if label is not None else {}), + **({"strict": strict} if strict is not None else {}), + }, + separators=(",", ":"), + ).encode("utf-8") + ) + if resolved_size > _MAX_GRAMMAR_BYTES: + raise _bad_request( + f"grammar payload {resolved_size} bytes exceeds limit ({_MAX_GRAMMAR_BYTES} bytes)", + param="grammar", + ) + else: + limit = _MAX_REGEX_LENGTH if kind == "regex" else _MAX_EBNF_LENGTH + if not isinstance(payload, str): + raise _bad_request(f"'grammar.{kind}' must be a string", param=f"grammar.{kind}") + if len(payload) > limit: + raise _bad_request(f"{kind} length {len(payload)} exceeds limit ({limit})", param=f"grammar.{kind}") + + return GrammarSpec(kind=cast("Any", kind), value=payload, label=label, strict=strict) + + +async def _render_native_image_prompt(config: Any, prompt: str, image_count: int) -> str: + """Render one image-aware user turn with the model's own chat template.""" + source = config.hf_id or config.weights_path + if not isinstance(source, str | Path): + raise _bad_request("model has no tokenizer source for image generation", param="images") + revision = config.hf_revision if config.hf_id else None + try: + tokenizer = await asyncio.to_thread( + _load_native_tokenizer_coalesced, + str(source), + revision, + ) + message = image_first_chat_message(role="user", text=prompt, image_count=image_count) + kwargs = dict(config.tasks.generate.chat_template_kwargs or {}) + apply_chat_template = cast("Any", tokenizer.apply_chat_template) + rendered = await asyncio.to_thread( + apply_chat_template, + [message], + tokenize=False, + add_generation_prompt=True, + **kwargs, + ) + except Exception as exc: + logger.info("native image prompt render failed for %s: %s", config.name, exc) + raise _bad_request("failed to render the model-native image prompt", param="images") from exc + if not isinstance(rendered, str) or not rendered: + raise _bad_request("model-native image prompt rendering returned no text", param="images") + return rendered + + +@lru_cache(maxsize=_NATIVE_TOKENIZER_CACHE_SIZE) +def _load_native_tokenizer_cached(source: str, revision: str | None) -> Any: + """Load one pinned tokenizer into the bounded direct-route cache.""" + return load_tokenizer(source, trust_remote_code=True, revision=revision) + + +def _load_native_tokenizer_coalesced(source: str, revision: str | None) -> Any: + """Coalesce concurrent cache misses without blocking the event loop.""" + with _NATIVE_TOKENIZER_LOAD_LOCK: + return _load_native_tokenizer_cached(source, revision) + + def _payload_too_large(message: str, *, param: str | None = None) -> HTTPException: """413 Payload Too Large, OpenAI-shaped error detail.""" detail: dict[str, Any] = { @@ -254,6 +583,25 @@ def _payload_too_large(message: str, *, param: str | None = None) -> HTTPExcepti return HTTPException(status_code=status.HTTP_413_CONTENT_TOO_LARGE, detail=detail) +async def _read_bounded_request_body(request: Request, limit: int) -> bytes: + """Read an ASGI request without aggregating more than ``limit`` bytes.""" + content_length = request.headers.get("content-length") + if content_length is not None: + try: + declared_length = int(content_length) + except ValueError: + declared_length = -1 + if declared_length > limit: + raise _payload_too_large(f"request body exceeds the limit of {limit} bytes") + + body = bytearray() + async for chunk in request.stream(): + if len(chunk) > limit - len(body): + raise _payload_too_large(f"request body exceeds the limit of {limit} bytes") + body.extend(chunk) + return bytes(body) + + async def _stream_generate_events( adapter: GenerationAdapter, *, @@ -266,10 +614,12 @@ async def _stream_generate_events( presence_penalty: float | None, top_k: int | None, min_new_tokens: int | None, + grammar: GrammarSpec | None, seed: int | None, logit_bias: dict[str, float] | None, logprobs: bool, top_logprobs: int | None, + images: list[ImageInput] | None = None, ) -> AsyncIterator[str]: """Yield SIE-native ``GenerateChunk`` SSE lines for ``SIEClient.stream_generate``. @@ -290,6 +640,11 @@ async def _stream_generate_events( completion_tokens = 0 saw_terminal = False terminal_error: dict[str, str] | None = None + optional_adapter_inputs: dict[str, Any] = {} + if grammar is not None: + optional_adapter_inputs["grammar"] = grammar + if images is not None: + optional_adapter_inputs["images"] = images try: async for chunk in adapter.generate( prompt=prompt, @@ -305,6 +660,7 @@ async def _stream_generate_events( logit_bias=logit_bias, logprobs=logprobs, top_logprobs=top_logprobs, + **optional_adapter_inputs, ): if chunk.done: saw_terminal = True @@ -473,8 +829,9 @@ async def generate( if x_machine_profile: span.set_attribute("machine_profile", x_machine_profile) + raw_body = await _read_bounded_request_body(http_request, _MAX_GENERATE_BODY_BYTES) try: - body = await http_request.json() + body = json.loads(raw_body) except (json.JSONDecodeError, ValueError) as exc: raise _bad_request("request body must be a JSON object") from exc if not isinstance(body, dict): @@ -509,6 +866,8 @@ async def generate( f"'prompt' is {prompt_bytes} bytes, exceeds the limit of {_MAX_PROMPT_BYTES} bytes", param="prompt", ) + images = _parse_native_images(body.get("images")) + grammar = _parse_native_grammar(body.get("grammar")) max_new_tokens = body.get("max_new_tokens") # ``isinstance(x, int)`` is True for ``bool`` in Python — reject @@ -537,6 +896,18 @@ async def generate( f"Model '{model}' does not declare a generate task", code=ErrorCode.MODEL_NOT_FOUND.value, ) + if images and not config.inputs.image: + raise _bad_request( + f"Model '{model}' does not support image input", + param="images", + code="unsupported_field", + ) + if grammar is not None and grammar.kind not in gen_task.capabilities.grammar: + raise _bad_request( + f"Model '{model}' does not declare '{grammar.kind}' grammar support", + param=f"grammar.{grammar.kind}", + code="unsupported_field", + ) if max_new_tokens > gen_task.max_output_tokens: raise _bad_request( f"max_new_tokens ({max_new_tokens}) exceeds model cap ({gen_task.max_output_tokens})", @@ -618,6 +989,10 @@ async def generate( ) logprobs, top_logprobs = _validate_logprobs(body.get("logprobs"), body.get("top_logprobs")) + generation_prompt = prompt + if images: + generation_prompt = await _render_native_image_prompt(config, prompt, len(images)) + # Do not start a potentially expensive model load until the complete # request has passed validation. await checker.ensure_loaded(device) @@ -633,7 +1008,7 @@ async def generate( return StreamingResponse( _stream_generate_events( adapter, - prompt=prompt, + prompt=generation_prompt, max_new_tokens=max_new_tokens, temperature=temperature, top_p=top_p, @@ -644,8 +1019,10 @@ async def generate( logit_bias=logit_bias, top_k=top_k, min_new_tokens=min_new_tokens, + grammar=grammar, logprobs=logprobs, top_logprobs=top_logprobs, + images=images, ), media_type="text/event-stream", headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}, @@ -656,8 +1033,13 @@ async def generate( # local-dev route keeps the walking-skeleton's blocking response shape # for backwards compatibility — drain the iterator into an # aggregate. SDK / gateway consume the iterator directly. + optional_adapter_inputs: dict[str, Any] = {} + if grammar is not None: + optional_adapter_inputs["grammar"] = grammar + if images is not None: + optional_adapter_inputs["images"] = images chunks = adapter.generate( - prompt=prompt, + prompt=generation_prompt, max_new_tokens=max_new_tokens, temperature=temperature, top_p=top_p, @@ -668,6 +1050,7 @@ async def generate( min_new_tokens=min_new_tokens, seed=seed, logit_bias=logit_bias, + **optional_adapter_inputs, ) result = await collect_generation(chunks) except Exception as e: diff --git a/packages/sie_server/src/sie_server/core/tokenizer.py b/packages/sie_server/src/sie_server/core/tokenizer.py index a152da1c..89ec36f5 100644 --- a/packages/sie_server/src/sie_server/core/tokenizer.py +++ b/packages/sie_server/src/sie_server/core/tokenizer.py @@ -3,7 +3,7 @@ from __future__ import annotations import logging -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any if TYPE_CHECKING: from pathlib import Path @@ -13,6 +13,19 @@ logger = logging.getLogger(__name__) +def image_first_chat_message(*, role: str, text: str, image_count: int) -> dict[str, Any]: + """Build the canonical image-first chat-template message. + + Native image ``generate`` and queue-backed message generation both use + this helper, keeping their model-template input byte-for-byte equivalent. + Image bytes travel separately; this value contains placeholders only. + """ + content: list[dict[str, str]] = [{"type": "image"} for _ in range(image_count)] + if text: + content.append({"type": "text", "text": text}) + return {"role": role, "content": content} + + def load_tokenizer( model_path: str | Path, *, diff --git a/packages/sie_server/src/sie_server/processors/streaming.py b/packages/sie_server/src/sie_server/processors/streaming.py index d6541e76..e7e76f4b 100644 --- a/packages/sie_server/src/sie_server/processors/streaming.py +++ b/packages/sie_server/src/sie_server/processors/streaming.py @@ -57,7 +57,7 @@ ) from sie_server.core.runtime_options import apply_generation_runtime_options from sie_server.core.text_tokens import estimate_tokens_from_chars -from sie_server.core.tokenizer import load_tokenizer +from sie_server.core.tokenizer import image_first_chat_message, load_tokenizer from sie_server.observability import worker_telemetry as _metrics from sie_server.processors.grammar_cache import GrammarLRU from sie_server.processors.grammar_compile import compile_outlines @@ -3219,10 +3219,7 @@ async def _render_chat_template( # key off the ``"image"`` part type; the bytes travel separately # as ``image_data``). Images lead, then the text — the flat image # order the adapter receives. - content_parts: list[dict[str, Any]] = [{"type": "image"} for _ in m.images] - if m.content: - content_parts.append({"type": "text", "text": m.content}) - d = {"role": role, "content": content_parts} + d = image_first_chat_message(role=role, text=m.content, image_count=len(m.images)) else: d = {"role": role, "content": m.content} if m.tool_calls: diff --git a/packages/sie_server/src/sie_server/types/openapi.py b/packages/sie_server/src/sie_server/types/openapi.py index aa5ead3a..2a10e239 100644 --- a/packages/sie_server/src/sie_server/types/openapi.py +++ b/packages/sie_server/src/sie_server/types/openapi.py @@ -270,10 +270,68 @@ class ScoreResponseModel(BaseModel): # Generate endpoint models +class NativeGenerateImageModel(BaseModel): + """One inline image on the SIE-native generate surface.""" + + data: str = Field( + ..., + min_length=1, + max_length=22_369_624, + description="Canonical standard-base64 image bytes, at most 16 MiB decoded", + ) + format: str | None = Field( + default=None, + min_length=1, + max_length=32, + pattern=r"^[A-Za-z0-9.+-]+$", + description="Short media-format hint such as png or jpeg", + ) + + +class NativeJsonSchemaGrammarModel(BaseModel): + """JSON Schema structured-output constraint.""" + + json_schema: dict[str, Any] + label: str | None = None + strict: bool | None = None + model_config = {"extra": "forbid"} + + +class NativeRegexGrammarModel(BaseModel): + """Regular-expression structured-output constraint.""" + + regex: str + label: str | None = None + strict: bool | None = None + model_config = {"extra": "forbid"} + + +class NativeEbnfGrammarModel(BaseModel): + """EBNF structured-output constraint.""" + + ebnf: str + label: str | None = None + strict: bool | None = None + model_config = {"extra": "forbid"} + + class GenerateRequestModel(BaseModel): """Request body for the SIE-native generate endpoint.""" prompt: str = Field(..., min_length=1, description="Prompt text to generate from") + images: list[NativeGenerateImageModel] | None = Field( + default=None, + min_length=1, + max_length=16, + description=( + "Optional inline images paired with prompt. The worker renders one user turn through " + "the model's native chat template; remote URLs are not accepted." + ), + ) + grammar: NativeJsonSchemaGrammarModel | NativeRegexGrammarModel | NativeEbnfGrammarModel | None = Field( + default=None, + description="Optional structured-output grammar; exactly one of json_schema, regex, or ebnf", + ) max_new_tokens: int = Field(..., ge=1, description="Maximum number of tokens to generate") temperature: float | None = Field(default=None, ge=0, description="Sampling temperature override") top_p: float | None = Field(default=None, gt=0, le=1, description="Nucleus-sampling probability override") diff --git a/packages/sie_server/tests/api/test_generate.py b/packages/sie_server/tests/api/test_generate.py index e386abb7..19fa85aa 100644 --- a/packages/sie_server/tests/api/test_generate.py +++ b/packages/sie_server/tests/api/test_generate.py @@ -8,25 +8,32 @@ from __future__ import annotations +import asyncio from collections.abc import AsyncIterator -from unittest.mock import MagicMock +from typing import Literal +from unittest.mock import AsyncMock, MagicMock import pytest from fastapi import FastAPI from fastapi.testclient import TestClient +from sie_sdk import SIEClient from sie_server.adapters._generation_base import GenerationAdapter, GenerationChunk from sie_server.adapters._spec import AdapterSpec from sie_server.adapters.base import ModelCapabilities, ModelDims +from sie_server.api import generate as generate_api from sie_server.api.generate import router as generate_router from sie_server.config.model import ( AdapterOptions, GenerateCapabilities, GenerateTask, + InputModalities, ModelConfig, ProfileConfig, Tasks, ) from sie_server.core.registry import ModelRegistry +from sie_server.types.grammar import GrammarSpec +from sie_server.types.inputs import ImageInput class _FakeGenAdapter(GenerationAdapter): @@ -67,10 +74,12 @@ async def generate( top_k: int | None = None, repetition_penalty: float | None = None, min_new_tokens: int | None = None, + grammar: GrammarSpec | None = None, seed: int | None = None, logit_bias: dict[str, float] | None = None, logprobs: bool = False, top_logprobs: int | None = None, + images: list[ImageInput] | None = None, ) -> AsyncIterator[GenerationChunk]: self.last_call = { "prompt": prompt, @@ -83,11 +92,14 @@ async def generate( "top_k": top_k, "repetition_penalty": repetition_penalty, "min_new_tokens": min_new_tokens, + "grammar": grammar, "seed": seed, "logit_bias": logit_bias, "logprobs": logprobs, "top_logprobs": top_logprobs, } + if images is not None: + self.last_call["images"] = images # Yield one delta + a terminal chunk so the local-dev route can # drain the iterator into the walking-skeleton-shaped aggregate response. yield GenerationChunk(text_delta=f"echo:{prompt}", is_first=True) @@ -100,7 +112,57 @@ async def generate( ) -def _make_config() -> ModelConfig: +class _LegacyTextGenAdapter(_FakeGenAdapter): + """Third-party-style adapter implementing the pre-grammar call signature.""" + + async def generate( + self, + prompt: str, + *, + max_new_tokens: int, + temperature: float = 1.0, + top_p: float = 1.0, + stop: list[str] | None = None, + frequency_penalty: float | None = None, + presence_penalty: float | None = None, + top_k: int | None = None, + repetition_penalty: float | None = None, + min_new_tokens: int | None = None, + seed: int | None = None, + logit_bias: dict[str, float] | None = None, + logprobs: bool = False, + top_logprobs: int | None = None, + ) -> AsyncIterator[GenerationChunk]: + self.last_call = { + "prompt": prompt, + "max_new_tokens": max_new_tokens, + "temperature": temperature, + "top_p": top_p, + "stop": stop, + "frequency_penalty": frequency_penalty, + "presence_penalty": presence_penalty, + "top_k": top_k, + "repetition_penalty": repetition_penalty, + "min_new_tokens": min_new_tokens, + "seed": seed, + "logit_bias": logit_bias, + "logprobs": logprobs, + "top_logprobs": top_logprobs, + } + yield GenerationChunk(text_delta=f"echo:{prompt}", is_first=True) + yield GenerationChunk( + text_delta="", + done=True, + finish_reason="stop", + prompt_tokens=len(prompt.split()), + completion_tokens=2, + ) + + +def _make_config( + *, + grammar: list[Literal["json_schema", "regex", "ebnf"]] | None = None, +) -> ModelConfig: return ModelConfig( sie_id="Qwen/Qwen3-4B-Instruct", hf_id="Qwen/Qwen3-4B-Instruct", @@ -108,7 +170,7 @@ def _make_config() -> ModelConfig: generate=GenerateTask( context_length=32768, max_output_tokens=4096, - capabilities=GenerateCapabilities(), + capabilities=GenerateCapabilities(grammar=grammar or []), ), ), profiles={ @@ -152,6 +214,394 @@ def client(registry: MagicMock) -> TestClient: class TestGenerateEndpoint: + @pytest.mark.parametrize("stream", [False, True]) + def test_text_only_generate_preserves_legacy_adapter_call_signature( + self, + client: TestClient, + registry: MagicMock, + stream: bool, + ) -> None: + legacy_adapter = _LegacyTextGenAdapter() + registry.get.return_value = legacy_adapter + + response = client.post( + "/v1/generate/Qwen__Qwen3-4B-Instruct", + json={"prompt": "Hello", "max_new_tokens": 8, "stream": stream}, + ) + + assert response.status_code == 200, response.text + assert legacy_adapter.last_call is not None + assert legacy_adapter.last_call["prompt"] == "Hello" + if stream: + assert '"finish_reason": "error"' not in response.text + assert "data: [DONE]" in response.text + + def test_sdk_multimodal_grammar_snippet_runs_against_local_route( + self, + client: TestClient, + registry: MagicMock, + fake_adapter: _FakeGenAdapter, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + config = _make_config(grammar=["json_schema"]) + config.inputs = InputModalities(text=True, image=True) + registry.get_config.return_value = config + + async def render(_config: object, prompt: str, image_count: int) -> str: + assert prompt == "Extract the title" + assert image_count == 1 + return "Extract the title" + + monkeypatch.setattr("sie_server.api.generate._render_native_image_prompt", render) + sdk = SIEClient("http://testserver") + sdk._client.close() + sdk._client = client # type: ignore[assignment] + + result = sdk.generate( + "Qwen/Qwen3-4B-Instruct", + "Extract the title", + images=[{"data": b"\x89PNG\r\n\x1a\npayload"}], + grammar={ + "json_schema": { + "type": "object", + "properties": {"title": {"type": "string"}}, + "required": ["title"], + } + }, + max_new_tokens=32, + ) + + assert result["model"] == "Qwen/Qwen3-4B-Instruct" + assert fake_adapter.last_call is not None + assert fake_adapter.last_call["prompt"] == "Extract the title" + assert fake_adapter.last_call["images"] == [{"data": b"\x89PNG\r\n\x1a\npayload", "format": "png"}] + assert fake_adapter.last_call["grammar"] == GrammarSpec( + kind="json_schema", + value={ + "type": "object", + "properties": {"title": {"type": "string"}}, + "required": ["title"], + }, + ) + + def test_request_body_is_rejected_before_unbounded_aggregation( + self, + client: TestClient, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + monkeypatch.setattr(generate_api, "_MAX_GENERATE_BODY_BYTES", 64) + + response = client.post( + "/v1/generate/Qwen__Qwen3-4B-Instruct", + json={"prompt": "x" * 128, "max_new_tokens": 8}, + ) + + assert response.status_code == 413 + assert response.json()["detail"]["code"] == "INPUT_TOO_LONG" + + @pytest.mark.parametrize("stream", [False, True]) + def test_native_images_render_model_prompt_and_reach_adapter_as_bytes( + self, + client: TestClient, + registry: MagicMock, + fake_adapter: _FakeGenAdapter, + monkeypatch: pytest.MonkeyPatch, + stream: bool, + ) -> None: + config = _make_config() + config.inputs = InputModalities(text=True, image=True) + registry.get_config.return_value = config + + async def render(_config: object, prompt: str, image_count: int) -> str: + assert prompt == "Read the image" + assert image_count == 1 + return "Read the image" + + monkeypatch.setattr("sie_server.api.generate._render_native_image_prompt", render) + response = client.post( + "/v1/generate/Qwen__Qwen3-4B-Instruct", + json={ + "prompt": "Read the image", + "images": [{"data": "aGVsbG8=", "format": "PNG"}], + "max_new_tokens": 8, + "stream": stream, + }, + ) + + assert response.status_code == 200, response.text + assert fake_adapter.last_call is not None + assert fake_adapter.last_call["prompt"] == "Read the image" + assert fake_adapter.last_call["images"] == [{"data": b"hello", "format": "png"}] + + @pytest.mark.parametrize( + ("grammar_body", "expected"), + [ + ( + { + "json_schema": { + "type": "object", + "properties": {"title": {"type": "string"}}, + "required": ["title"], + }, + "label": "document", + "strict": True, + }, + GrammarSpec( + kind="json_schema", + value={ + "type": "object", + "properties": {"title": {"type": "string"}}, + "required": ["title"], + }, + label="document", + strict=True, + ), + ), + ({"regex": r"\d{3}-\d{4}"}, GrammarSpec(kind="regex", value=r"\d{3}-\d{4}")), + ({"ebnf": 'root ::= "yes" | "no"'}, GrammarSpec(kind="ebnf", value='root ::= "yes" | "no"')), + ], + ) + @pytest.mark.parametrize("stream", [False, True]) + def test_native_grammar_reaches_adapter( + self, + client: TestClient, + registry: MagicMock, + fake_adapter: _FakeGenAdapter, + grammar_body: dict[str, object], + expected: GrammarSpec, + stream: bool, + ) -> None: + registry.get_config.return_value = _make_config(grammar=[expected.kind]) + + response = client.post( + "/v1/generate/Qwen__Qwen3-4B-Instruct", + json={ + "prompt": "Return structured output", + "grammar": grammar_body, + "max_new_tokens": 8, + "stream": stream, + }, + ) + + assert response.status_code == 200, response.text + assert fake_adapter.last_call is not None + assert fake_adapter.last_call["grammar"] == expected + + def test_native_grammar_dereferences_internal_schema_refs( + self, + client: TestClient, + registry: MagicMock, + fake_adapter: _FakeGenAdapter, + ) -> None: + registry.get_config.return_value = _make_config(grammar=["json_schema"]) + + response = client.post( + "/v1/generate/Qwen__Qwen3-4B-Instruct", + json={ + "prompt": "Return structured output", + "grammar": { + "json_schema": { + "$defs": {"Title": {"type": "string", "minLength": 1}}, + "type": "object", + "properties": {"title": {"$ref": "#/$defs/Title"}}, + } + }, + "max_new_tokens": 8, + }, + ) + + assert response.status_code == 200, response.text + assert fake_adapter.last_call is not None + grammar = fake_adapter.last_call["grammar"] + assert isinstance(grammar, GrammarSpec) + assert grammar.value == { + "type": "object", + "properties": {"title": {"type": "string", "minLength": 1}}, + } + + @pytest.mark.parametrize( + "grammar", + [ + {}, + {"json_schema": {}, "regex": "x"}, + {"regex": 123}, + {"ebnf": "root", "unknown": True}, + {"json_schema": {"$ref": "https://example.com/schema.json"}}, + ], + ) + def test_native_grammar_rejects_invalid_shape_before_adapter( + self, + client: TestClient, + registry: MagicMock, + fake_adapter: _FakeGenAdapter, + grammar: dict[str, object], + ) -> None: + registry.get_config.return_value = _make_config(grammar=["json_schema", "regex", "ebnf"]) + + response = client.post( + "/v1/generate/Qwen__Qwen3-4B-Instruct", + json={"prompt": "Hi", "grammar": grammar, "max_new_tokens": 8}, + ) + + assert response.status_code == 400 + assert fake_adapter.last_call is None + + def test_native_grammar_requires_model_capability( + self, + client: TestClient, + fake_adapter: _FakeGenAdapter, + ) -> None: + response = client.post( + "/v1/generate/Qwen__Qwen3-4B-Instruct", + json={"prompt": "Hi", "grammar": {"regex": "[a-z]+"}, "max_new_tokens": 8}, + ) + + assert response.status_code == 400 + assert response.json()["detail"]["param"] == "grammar.regex" + assert fake_adapter.last_call is None + + def test_native_image_prompt_uses_pinned_trusted_model_tokenizer( + self, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + config = _make_config() + config.hf_revision = "0123456789abcdef0123456789abcdef01234567" + assert config.tasks.generate is not None + config.tasks.generate.chat_template_kwargs = {"enable_thinking": False} + captured: dict[str, object] = {} + + class _Tokenizer: + def apply_chat_template(self, messages: object, **kwargs: object) -> str: + captured["messages"] = messages + captured["template_kwargs"] = kwargs + return "Read" + + def load(source: object, **kwargs: object) -> _Tokenizer: + captured["source"] = source + captured["load_kwargs"] = kwargs + return _Tokenizer() + + generate_api._load_native_tokenizer_cached.cache_clear() + monkeypatch.setattr(generate_api, "load_tokenizer", load) + rendered = asyncio.run(generate_api._render_native_image_prompt(config, "Read", 1)) + + assert rendered == "Read" + assert captured["source"] == "Qwen/Qwen3-4B-Instruct" + assert captured["load_kwargs"] == { + "trust_remote_code": True, + "revision": "0123456789abcdef0123456789abcdef01234567", + } + assert captured["messages"] == [ + {"role": "user", "content": [{"type": "image"}, {"type": "text", "text": "Read"}]} + ] + assert captured["template_kwargs"] == { + "tokenize": False, + "add_generation_prompt": True, + "enable_thinking": False, + } + + def test_native_image_prompt_coalesces_and_caches_tokenizer_loads( + self, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + config = _make_config() + config.hf_revision = "0123456789abcdef0123456789abcdef01234567" + loads = 0 + + class _Tokenizer: + def apply_chat_template(self, _messages: object, **_kwargs: object) -> str: + return "Read" + + def load(_source: object, **_kwargs: object) -> _Tokenizer: + nonlocal loads + loads += 1 + return _Tokenizer() + + async def render_twice() -> list[str]: + return await asyncio.gather( + generate_api._render_native_image_prompt(config, "Read", 1), + generate_api._render_native_image_prompt(config, "Read", 1), + ) + + generate_api._load_native_tokenizer_cached.cache_clear() + monkeypatch.setattr(generate_api, "load_tokenizer", load) + rendered = asyncio.run(render_twice()) + + assert rendered == ["Read", "Read"] + assert loads == 1 + + @pytest.mark.parametrize( + ("images", "expected_param"), + [ + ([], "images"), + ([{"data": "!!!"}], "images[0].data"), + ([{"data": "aGk"}], "images[0].data"), + ([{"data": "__8="}], "images[0].data"), + ([{"data": "AB=="}], "images[0].data"), + ([{"data": "aGVsbG8=", "url": "https://example.com/a.png"}], "images[0].url"), + ([{"data": "aGVsbG8=", "format": "png;bad"}], "images[0].format"), + ([{"data": "aGk="}] * 17, "images"), + ], + ) + def test_native_images_reject_malformed_envelopes_before_load( + self, + client: TestClient, + registry: MagicMock, + images: object, + expected_param: str, + ) -> None: + response = client.post( + "/v1/generate/Qwen__Qwen3-4B-Instruct", + json={"prompt": "Read", "images": images, "max_new_tokens": 8}, + ) + + assert response.status_code == 400 + assert response.json()["detail"]["param"] == expected_param + registry.load_async.assert_not_called() + + def test_native_images_reject_nonvision_model_before_load( + self, + client: TestClient, + registry: MagicMock, + ) -> None: + response = client.post( + "/v1/generate/Qwen__Qwen3-4B-Instruct", + json={"prompt": "Read", "images": [{"data": "aGVsbG8="}], "max_new_tokens": 8}, + ) + + assert response.status_code == 400 + assert response.json()["detail"] == { + "code": "unsupported_field", + "message": "Model 'Qwen__Qwen3-4B-Instruct' does not support image input", + "param": "images", + } + registry.load_async.assert_not_called() + + def test_native_image_template_failure_is_rejected_before_model_load( + self, + client: TestClient, + registry: MagicMock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + config = _make_config() + config.inputs = InputModalities(text=True, image=True) + registry.get_config.return_value = config + registry.is_loaded.return_value = False + registry.load_async = AsyncMock() + + async def reject_template(_config: object, _prompt: str, _image_count: int) -> str: + raise generate_api._bad_request("image prompt template failed", param="images") + + monkeypatch.setattr(generate_api, "_render_native_image_prompt", reject_template) + response = client.post( + "/v1/generate/Qwen__Qwen3-4B-Instruct", + json={"prompt": "Read", "images": [{"data": "aGVsbG8="}], "max_new_tokens": 8}, + ) + + assert response.status_code == 400 + assert response.json()["detail"]["param"] == "images" + registry.load_async.assert_not_awaited() + def test_happy_path_returns_text_finish_reason_usage( self, client: TestClient, fake_adapter: _FakeGenAdapter ) -> None: @@ -181,6 +631,7 @@ def test_happy_path_returns_text_finish_reason_usage( "top_k": None, "repetition_penalty": None, "min_new_tokens": None, + "grammar": None, "seed": None, "logit_bias": None, "logprobs": False, @@ -436,26 +887,6 @@ def test_valid_penalty_is_forwarded(self, client: TestClient, fake_adapter: _Fak assert fake_adapter.last_call is not None assert fake_adapter.last_call[field] == 0.5 - @pytest.mark.parametrize( - "grammar", - [ - "not-an-object", - {"regex": "[a-z]+"}, - {"regex": "[a-z]+", "ebnf": 'root ::= "x"'}, - ], - ) - def test_grammar_is_rejected_when_adapter_cannot_apply_it(self, client: TestClient, grammar: object) -> None: - response = client.post( - "/v1/generate/Qwen__Qwen3-4B-Instruct", - json={"prompt": "Hi", "max_new_tokens": 8, "grammar": grammar}, - ) - assert response.status_code == 400, response.text - assert response.json()["detail"] == { - "code": "unsupported_field", - "message": "unsupported field(s): ['grammar']", - "param": "grammar", - } - def test_prompt_at_cap_is_accepted(self, client: TestClient, monkeypatch: pytest.MonkeyPatch) -> None: # A prompt exactly at the byte cap is allowed (boundary check). monkeypatch.setattr("sie_server.api.generate._MAX_PROMPT_BYTES", 16) diff --git a/packages/sie_server/tests/test_openapi_export.py b/packages/sie_server/tests/test_openapi_export.py index 1b441a8d..8306bf1e 100644 --- a/packages/sie_server/tests/test_openapi_export.py +++ b/packages/sie_server/tests/test_openapi_export.py @@ -51,6 +51,19 @@ def test_openapi_documents_generate_contract() -> None: assert request_schema == {"$ref": "#/components/schemas/GenerateRequestModel"} schema = spec["components"]["schemas"]["GenerateRequestModel"] + images = schema["properties"]["images"] + assert images["anyOf"][0]["items"]["$ref"] == "#/components/schemas/NativeGenerateImageModel" + assert images["anyOf"][0]["minItems"] == 1 + assert images["anyOf"][0]["maxItems"] == 16 + image_schema = spec["components"]["schemas"]["NativeGenerateImageModel"] + assert image_schema["properties"]["data"]["maxLength"] == 22_369_624 + grammar = schema["properties"]["grammar"] + grammar_refs = {variant["$ref"] for variant in grammar["anyOf"] if "$ref" in variant} + assert grammar_refs == { + "#/components/schemas/NativeJsonSchemaGrammarModel", + "#/components/schemas/NativeRegexGrammarModel", + "#/components/schemas/NativeEbnfGrammarModel", + } assert set(schema["required"]) == {"prompt", "max_new_tokens"} assert schema["properties"]["stream"]["anyOf"][0] == {"type": "boolean"} seed_schema = schema["properties"]["seed"] @@ -61,7 +74,7 @@ def test_openapi_documents_generate_contract() -> None: assert schema["properties"]["logprobs"]["anyOf"][0] == {"type": "boolean"} assert schema["properties"]["top_logprobs"]["anyOf"][0]["minimum"] == 0 assert schema["properties"]["top_logprobs"]["anyOf"][0]["maximum"] == 20 - for unsupported in ("grammar", "lora_adapter", "n", "best_of", "stream_options"): + for unsupported in ("lora_adapter", "n", "best_of", "stream_options"): assert unsupported not in schema["properties"] response_content = operation["responses"]["200"]["content"] diff --git a/packages/sie_ts_sdk/package.json b/packages/sie_ts_sdk/package.json index 1b101aff..1bfe48e7 100644 --- a/packages/sie_ts_sdk/package.json +++ b/packages/sie_ts_sdk/package.json @@ -34,7 +34,7 @@ ], "scripts": { "build": "tsup", - "typecheck": "tsc --noEmit", + "typecheck": "tsc --noEmit && tsc -p tsconfig.type-tests.json", "lint": "biome check .", "lint:fix": "biome check --write .", "test": "vitest run", diff --git a/packages/sie_ts_sdk/src/client.ts b/packages/sie_ts_sdk/src/client.ts index bae903e3..51113f0b 100644 --- a/packages/sie_ts_sdk/src/client.ts +++ b/packages/sie_ts_sdk/src/client.ts @@ -39,7 +39,7 @@ import { SIEStreamError, ServerError, } from "./errors.js"; -import { toImageWireFormat } from "./images.js"; +import { detectImageFormat, toImageBytes, toImageWireFormat } from "./images.js"; import type { ImageInput, ImageWireFormat } from "./images.js"; import { DEFAULT_JOB_WAIT_POLL, @@ -109,6 +109,8 @@ import type { ExtractResult, FileDeleted, GenerateChunk, + GenerateGrammar, + GenerateImage, GenerateOptions, GenerateResult, Item, @@ -227,10 +229,7 @@ function parseRequestMetadata(headers: Headers): RequestMetadata | undefined { metadata.id = requestId; } const executionIdentitySha256 = headers.get("x-sie-execution-identity-sha256"); - if ( - executionIdentitySha256 !== null && - /^[0-9a-f]{64}$/.test(executionIdentitySha256) - ) { + if (executionIdentitySha256 !== null && /^[0-9a-f]{64}$/.test(executionIdentitySha256)) { metadata.executionIdentitySha256 = executionIdentitySha256; } @@ -277,6 +276,41 @@ function validateGenerationSeed(seed: number): number { return seed; } +/** Validate the discriminated native grammar before issuing a billable request. */ +function validateGenerateGrammar(grammar: GenerateGrammar | Record): void { + if (typeof grammar !== "object" || grammar === null || Array.isArray(grammar)) { + throw new TypeError("grammar must be an object"); + } + const allowed = new Set(["json_schema", "regex", "ebnf", "label", "strict"]); + const unknown = Object.keys(grammar).filter((key) => !allowed.has(key)); + if (unknown.length > 0) { + throw new TypeError(`grammar contains unsupported field(s): ${unknown.sort().join(", ")}`); + } + const variants = ["json_schema", "regex", "ebnf"].filter((key) => Object.hasOwn(grammar, key)); + if (variants.length !== 1) { + throw new TypeError("grammar must contain exactly one of json_schema, regex, or ebnf"); + } + const variant = variants[0]; + const value = grammar[variant as keyof GenerateGrammar]; + if ( + variant === "json_schema" + ? typeof value !== "object" || value === null || Array.isArray(value) + : typeof value !== "string" + ) { + throw new TypeError( + variant === "json_schema" + ? "grammar.json_schema must be an object" + : `grammar.${variant} must be a string`, + ); + } + if (grammar.label !== undefined && typeof grammar.label !== "string") { + throw new TypeError("grammar.label must be a string"); + } + if (grammar.strict !== undefined && typeof grammar.strict !== "boolean") { + throw new TypeError("grammar.strict must be a boolean"); + } +} + /** Serialize the controls shared by blocking and streaming native generation. */ function applyGenerateOptions(body: Record, options: GenerateOptions): void { if (options.temperature !== undefined) body.temperature = options.temperature; @@ -285,7 +319,10 @@ function applyGenerateOptions(body: Record, options: GenerateOp if (options.stop !== undefined) body.stop = options.stop; if (options.frequencyPenalty !== undefined) body.frequency_penalty = options.frequencyPenalty; if (options.presencePenalty !== undefined) body.presence_penalty = options.presencePenalty; - if (options.grammar !== undefined) body.grammar = options.grammar; + if (options.grammar !== undefined) { + validateGenerateGrammar(options.grammar); + body.grammar = options.grammar; + } if (options.seed !== undefined) body.seed = validateGenerationSeed(options.seed); if (options.logitBias !== undefined) body.logit_bias = options.logitBias; if (options.routingKey !== undefined) body.routing_key = options.routingKey; @@ -348,6 +385,59 @@ async function imageForWire(image: ImageInput | ImageWireFormat): Promise { + if (images.length === 0 || images.length > 16) { + throw new RangeError("images must contain between 1 and 16 entries"); + } + return Promise.all( + images.map(async (image) => { + const data = await toImageBytes(isGenerateImage(image) ? image.data : image); + if (data.byteLength === 0 || data.byteLength > 16 * 1024 * 1024) { + throw new RangeError("each image must contain between 1 byte and 16 MiB"); + } + const detected = detectImageFormat(data); + const declared = isGenerateImage(image) ? image.format : undefined; + let format: string; + if (declared === undefined) { + if (detected === "unknown") { + throw new TypeError("could not detect image format; provide a format hint"); + } + format = detected; + } else { + if (!/^[A-Za-z0-9.+-]{1,32}$/.test(declared)) { + throw new TypeError("image format must be a short ASCII media-format token"); + } + format = declared.toLowerCase(); + if (format === "jpg" || format === "jpe") format = "jpeg"; + if (detected !== "unknown" && format !== detected) { + throw new TypeError( + `image format mismatch: declared '${format}', detected '${detected}'`, + ); + } + } + return { data: imageBytesToBase64(data), format }; + }), + ); +} + async function itemImagesForWire(item: Item): Promise { if (!item.images || item.images.length === 0) { return item; @@ -821,12 +911,30 @@ export class SIEClient { * console.log(result.text); * console.log(`TTFT: ${result.ttftMs}ms`); * ``` + * + * Vision and structured output use the same method: + * ```typescript + * const result = await client.generate("vision-model", "Extract the title.", { + * maxNewTokens: 128, + * images: [{ data: imageBytes, format: "png" }], + * grammar: { + * json_schema: { + * type: "object", + * properties: { title: { type: "string" } }, + * required: ["title"], + * }, + * }, + * }); + * ``` */ async generate(model: string, prompt: string, options: GenerateOptions): Promise { const body: Record = { prompt, max_new_tokens: options.maxNewTokens, }; + if (options.images !== undefined) { + body.images = await generationImagesForWire(options.images); + } applyGenerateOptions(body, options); const { pool, gpu } = this.parseGpuParam(options.gpu); @@ -1085,6 +1193,9 @@ export class SIEClient { max_new_tokens: options.maxNewTokens, stream: true, }; + if (options.images !== undefined) { + body.images = await generationImagesForWire(options.images); + } applyGenerateOptions(body, options); if (options.logprobs !== undefined) body.logprobs = options.logprobs; if (options.topLogprobs !== undefined) { diff --git a/packages/sie_ts_sdk/src/index.ts b/packages/sie_ts_sdk/src/index.ts index 84accff9..9d41d197 100644 --- a/packages/sie_ts_sdk/src/index.ts +++ b/packages/sie_ts_sdk/src/index.ts @@ -109,6 +109,11 @@ export type { // Generation FinishReason, GenerationUsage, + GenerateImage, + GenerateGrammar, + JsonSchemaGrammar, + RegexGrammar, + EbnfGrammar, GenerateOptions, StreamGenerateOptions, GenerateResult, diff --git a/packages/sie_ts_sdk/src/types.ts b/packages/sie_ts_sdk/src/types.ts index 65086f39..ac7259db 100644 --- a/packages/sie_ts_sdk/src/types.ts +++ b/packages/sie_ts_sdk/src/types.ts @@ -727,10 +727,51 @@ export interface GenerationUsage { totalTokens: number; } +/** One image paired with a native generation prompt. */ +export interface GenerateImage { + /** Encoded image bytes, a Blob/File, or base64/data-URL text. */ + data: ImageInput; + /** Optional media format hint. The SDK infers JPEG, PNG, and WebP. */ + format?: string; +} + +interface GrammarMetadata { + /** Optional schema/grammar label used by structured-output backends. */ + label?: string; + /** Optional strictness hint for structured-output backends. */ + strict?: boolean; +} + +/** Constrain native generation to a JSON Schema. */ +export type JsonSchemaGrammar = GrammarMetadata & { + json_schema: Record; + regex?: never; + ebnf?: never; +}; + +/** Constrain native generation to a regular expression. */ +export type RegexGrammar = GrammarMetadata & { + json_schema?: never; + regex: string; + ebnf?: never; +}; + +/** Constrain native generation to an EBNF grammar. */ +export type EbnfGrammar = GrammarMetadata & { + json_schema?: never; + regex?: never; + ebnf: string; +}; + +/** Native structured-output grammar. Exactly one grammar variant is set. */ +export type GenerateGrammar = JsonSchemaGrammar | RegexGrammar | EbnfGrammar; + /** Options for the generate operation. */ export interface GenerateOptions { /** Hard cap on output tokens. Required. */ maxNewTokens: number; + /** Optional images rendered with the prompt for vision-capable models. */ + images?: (ImageInput | GenerateImage)[]; /** Sampling temperature. */ temperature?: number; /** Nucleus sampling cutoff. */ @@ -741,8 +782,14 @@ export interface GenerateOptions { frequencyPenalty?: number; /** OpenAI-compatible presence penalty in [-2, 2]. */ presencePenalty?: number; - /** Native structured-output grammar. */ - grammar?: Record; + /** + * Native structured-output grammar. + * + * The broad record arm preserves the pre-existing SDK input contract for + * callers that keep valid grammar objects in `Record` + * variables. The SDK still validates the exact three-arm shape at runtime. + */ + grammar?: GenerateGrammar | Record; /** * Optional per-request sampling seed. Must be a JavaScript safe integer * (-(2^53 - 1) through 2^53 - 1) so JSON serialization preserves it exactly. diff --git a/packages/sie_ts_sdk/tests/generate.test.ts b/packages/sie_ts_sdk/tests/generate.test.ts index 1a1463ad..30c439a5 100644 --- a/packages/sie_ts_sdk/tests/generate.test.ts +++ b/packages/sie_ts_sdk/tests/generate.test.ts @@ -65,6 +65,7 @@ describe("SIEClient.generate", () => { const omittedSamplerBody = JSON.parse(mockFetch.mock.calls[0][1].body); expect(omittedSamplerBody).not.toHaveProperty("temperature"); expect(omittedSamplerBody).not.toHaveProperty("top_p"); + expect(mockFetch.mock.calls[0][1].body).toBe('{"prompt":"Hi","max_new_tokens":32}'); }); it("sends a JSON body with snake_case field names", async () => { @@ -122,6 +123,84 @@ describe("SIEClient.generate", () => { }); }); + it("accepts a valid grammar held in the legacy broad record type", async () => { + mockFetch.mockResolvedValueOnce( + jsonResponse({ + model: "m", + text: "123", + finish_reason: "stop", + usage: { prompt_tokens: 1, completion_tokens: 1, total_tokens: 2 }, + }), + ); + const grammar: Record = { regex: "\\d+" }; + + const client = new SIEClient("http://localhost:8080"); + await client.generate("m", "Return digits", { maxNewTokens: 8, grammar }); + + const body = JSON.parse(mockFetch.mock.calls[0][1].body); + expect(body.grammar).toEqual(grammar); + }); + + it("serializes images and typed JSON-schema grammar", async () => { + mockFetch.mockResolvedValueOnce( + jsonResponse({ + model: "vision-model", + text: '{"title":"Launch"}', + finish_reason: "stop", + usage: { prompt_tokens: 1, completion_tokens: 1, total_tokens: 2 }, + }), + ); + + const client = new SIEClient("http://localhost:8080"); + await client.generate("vision-model", "Read this document", { + maxNewTokens: 32, + images: [{ data: Uint8Array.from([0x89, 0x50, 0x4e, 0x47, 13, 10, 26, 10, 1]) }], + grammar: { + json_schema: { type: "object", properties: { title: { type: "string" } } }, + label: "document", + strict: true, + }, + }); + + const body = JSON.parse(mockFetch.mock.calls[0][1].body); + expect(body.images).toEqual([{ data: "iVBORw0KGgoB", format: "png" }]); + expect(body.grammar).toEqual({ + json_schema: { type: "object", properties: { title: { type: "string" } } }, + label: "document", + strict: true, + }); + }); + + it.each([{}, { json_schema: {}, regex: "x" }, { regex: 123 }, { ebnf: "root", unknown: true }])( + "rejects an invalid grammar before request: %j", + async (grammar) => { + const client = new SIEClient("http://localhost:8080"); + await expect( + client.generate("m", "Hi", { + maxNewTokens: 8, + grammar: grammar as never, + }), + ).rejects.toThrow(); + expect(mockFetch).not.toHaveBeenCalled(); + }, + ); + + it("rejects a mismatched image format before request", async () => { + const client = new SIEClient("http://localhost:8080"); + await expect( + client.generate("m", "Hi", { + maxNewTokens: 8, + images: [ + { + data: Uint8Array.from([0x89, 0x50, 0x4e, 0x47, 13, 10, 26, 10, 1]), + format: "jpeg", + }, + ], + }), + ).rejects.toThrow("image format mismatch"); + expect(mockFetch).not.toHaveBeenCalled(); + }); + it.each([Number.MAX_SAFE_INTEGER + 1, Number.MIN_SAFE_INTEGER - 1, 1.5, Number.NaN])( "rejects a seed that JSON cannot preserve exactly: %s", async (seed) => { diff --git a/packages/sie_ts_sdk/tests/streaming.test.ts b/packages/sie_ts_sdk/tests/streaming.test.ts index 89d0fd5c..7752e443 100644 --- a/packages/sie_ts_sdk/tests/streaming.test.ts +++ b/packages/sie_ts_sdk/tests/streaming.test.ts @@ -498,6 +498,7 @@ describe("SIEClient.streamGenerate", () => { const client = new SIEClient("http://localhost:8080"); for await (const _ of client.streamGenerate("Qwen/Qwen3-4B-Instruct-2507", "hi", { maxNewTokens: 4, + images: [Uint8Array.from([0xff, 0xd8, 0xff, 0xe0, 1])], temperature: 0.5, topP: 0.9, stop: [""], @@ -521,6 +522,7 @@ describe("SIEClient.streamGenerate", () => { expect(body).toEqual({ prompt: "hi", max_new_tokens: 4, + images: [{ data: "/9j/4AE=", format: "jpeg" }], temperature: 0.5, top_p: 0.9, stop: [""], diff --git a/packages/sie_ts_sdk/tests/typecheck/generate-grammar.ts b/packages/sie_ts_sdk/tests/typecheck/generate-grammar.ts new file mode 100644 index 00000000..8aba5db7 --- /dev/null +++ b/packages/sie_ts_sdk/tests/typecheck/generate-grammar.ts @@ -0,0 +1,16 @@ +import type { SIEClient } from "../../src/client.js"; +import type { GenerateOptions } from "../../src/types.js"; + +declare const client: SIEClient; + +// Existing callers commonly retain grammar configuration in a broad record. +// This must remain source-compatible while runtime validation enforces the +// native json_schema | regex | ebnf envelope. +const legacyGrammar: Record = { regex: "\\d+" }; +const options: GenerateOptions = { + maxNewTokens: 8, + grammar: legacyGrammar, +}; + +void client.generate("model", "Return digits", options); +void client.streamGenerate("model", "Return digits", options); diff --git a/packages/sie_ts_sdk/tsconfig.type-tests.json b/packages/sie_ts_sdk/tsconfig.type-tests.json new file mode 100644 index 00000000..5606574f --- /dev/null +++ b/packages/sie_ts_sdk/tsconfig.type-tests.json @@ -0,0 +1,9 @@ +{ + "extends": "./tsconfig.json", + "compilerOptions": { + "noEmit": true, + "rootDir": "." + }, + "include": ["src/**/*", "tests/typecheck/**/*.ts"], + "exclude": ["node_modules", "dist"] +}