Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
72 changes: 72 additions & 0 deletions packages/sie_gateway/openapi.json
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,36 @@
],
"type": "object"
},
"AudioInput": {
"properties": {
"data": {
"items": {
"format": "int32",
"minimum": 0,
"type": "integer"
},
"type": "array"
},
"format": {
"type": [
"string",
"null"
]
},
"sample_rate": {
"format": "int32",
"minimum": 0,
"type": [
"integer",
"null"
]
}
},
"required": [
"data"
],
"type": "object"
},
"BundleConfigDocument": {
"properties": {
"adapters": {
Expand Down Expand Up @@ -1702,6 +1732,16 @@
},
"ItemInput": {
"properties": {
"audio": {
"oneOf": [
{
"type": "null"
},
{
"$ref": "#/components/schemas/AudioInput"
}
]
},
"document": {
"oneOf": [
{
Expand Down Expand Up @@ -1733,6 +1773,16 @@
"string",
"null"
]
},
"video": {
"oneOf": [
{
"type": "null"
},
{
"$ref": "#/components/schemas/VideoInput"
}
]
}
},
"type": "object"
Expand Down Expand Up @@ -2930,6 +2980,28 @@
],
"type": "object"
},
"VideoInput": {
"properties": {
"data": {
"items": {
"format": "int32",
"minimum": 0,
"type": "integer"
},
"type": "array"
},
"format": {
"type": [
"string",
"null"
]
}
},
"required": [
"data"
],
"type": "object"
},
"WorkerInfo": {
"properties": {
"bundle": {
Expand Down
37 changes: 37 additions & 0 deletions packages/sie_gateway/src/openapi.rs
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,7 @@ static OPENAPI_JSON: LazyLock<String> = LazyLock::new(|| {
InferenceInternalServerErrorResponse,
InferenceServiceUnavailableResponse,
AllItemsFailedResponse,
AudioInput,
BundleRoutingConflictDetail,
BundleConflictResponse,
DocumentInput,
Expand Down Expand Up @@ -144,6 +145,7 @@ static OPENAPI_JSON: LazyLock<String> = LazyLock::new(|| {
ScoreResponse,
SparseVector,
TimingInfo,
VideoInput,
crate::types::pool::AssignedWorker,
crate::types::worker::WorkerInfo
)),
Expand Down Expand Up @@ -2173,6 +2175,22 @@ pub struct ImageInput {
pub format: Option<String>,
}

#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct AudioInput {
pub data: Vec<u8>,
#[serde(default)]
pub format: Option<String>,
#[serde(default)]
pub sample_rate: Option<u32>,
}

#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct VideoInput {
pub data: Vec<u8>,
#[serde(default)]
pub format: Option<String>,
}

#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct DocumentInput {
pub data: Vec<u8>,
Expand All @@ -2189,6 +2207,10 @@ pub struct ItemInput {
#[serde(default)]
pub images: Option<Vec<ImageInput>>,
#[serde(default)]
pub audio: Option<AudioInput>,
#[serde(default)]
pub video: Option<VideoInput>,
#[serde(default)]
pub document: Option<DocumentInput>,
#[serde(default)]
pub metadata: Option<Value>,
Expand Down Expand Up @@ -2417,6 +2439,21 @@ mod tests {
);
}

#[test]
fn openapi_json_documents_native_audio_and_video_inputs() {
let spec: serde_json::Value = serde_json::from_str(&OPENAPI_JSON).unwrap();
let properties = &spec["components"]["schemas"]["ItemInput"]["properties"];

assert_eq!(
properties["audio"]["oneOf"][1]["$ref"],
"#/components/schemas/AudioInput"
);
assert_eq!(
properties["video"]["oneOf"][1]["$ref"],
"#/components/schemas/VideoInput"
);
}

#[tokio::test]
async fn docs_ui_serves_self_contained_redoc_html() {
let resp = docs_ui().await.into_response();
Expand Down
21 changes: 14 additions & 7 deletions packages/sie_sdk/src/sie_sdk/client/_shared.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
import numpy as np

from sie_sdk.images import convert_item_images
from sie_sdk.media import convert_item_media

_logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -121,13 +122,19 @@ def json(self) -> Any: ...
_retry_rng = random.Random() # noqa: S311 — non-cryptographic jitter only


def convert_score_images_for_wire(query: Any, items: Sequence[Any]) -> tuple[Any, list[Any]]:
"""Convert image-bearing score query/items to the SDK image wire shape."""
query_for_wire = convert_item_images({**query}) if "images" in query else query
items_for_wire = [
convert_item_images({**item}) if "images" in item else item # ty: ignore[invalid-argument-type]
for item in items
]
def convert_score_media_for_wire(query: Any, items: Sequence[Any]) -> tuple[Any, list[Any]]:
"""Convert media-bearing score query/items to SDK wire shapes."""

def convert(item: Any) -> Any:
if not any(field in item for field in ("images", "audio", "video")):
return item
item_for_wire = {**item}
if "images" in item_for_wire:
item_for_wire = convert_item_images(item_for_wire)
return convert_item_media(item_for_wire)

query_for_wire = convert(query)
items_for_wire = [convert(item) for item in items]
return query_for_wire, items_for_wire


Expand Down
25 changes: 16 additions & 9 deletions packages/sie_sdk/src/sie_sdk/client/async_.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,7 @@
from sie_sdk.files import resolve_upload
from sie_sdk.images import convert_item_images
from sie_sdk.jobs import TERMINAL_JOB_STATES, build_job_body, decode_chunk_bytes, job_chunks
from sie_sdk.media import convert_item_media
from sie_sdk.types import (
Batch,
CapacityInfo,
Expand Down Expand Up @@ -97,7 +98,7 @@
check_version_skew,
compute_oom_backoff,
compute_retry_delay,
convert_score_images_for_wire,
convert_score_media_for_wire,
get_error_code,
get_retry_after,
get_sdk_version,
Expand Down Expand Up @@ -1131,12 +1132,17 @@ async def encode(
single_item = not isinstance(items, list)
items_list = [items] if single_item else items

# Convert images to JPEG bytes for transport.
# Only copy items that have images — text-only items are passed through directly
items_for_wire = [
convert_item_images({**item}) if "images" in item else item # ty: ignore[invalid-argument-type]
for item in items_list
]
# Convert media to bytes for transport.
# Only copy media-bearing items; text-only items pass through directly.
items_for_wire = []
for item in items_list:
if not any(field in item for field in ("images", "audio", "video")):
items_for_wire.append(item)
continue
wire_item: dict[str, Any] = {**item} # ty: ignore[invalid-argument-type]
if "images" in wire_item:
wire_item = convert_item_images(wire_item)
items_for_wire.append(convert_item_media(wire_item))

# Build request body
request_body: dict[str, Any] = {"items": items_for_wire}
Expand Down Expand Up @@ -1600,7 +1606,7 @@ async def score(
# Resolve defaults and pool
pool_name, resolved_gpu = await self._resolve_pool_and_gpu(gpu)
resolved_options = self._resolve_options(options)
query_for_wire, items_for_wire = convert_score_images_for_wire(query, items)
query_for_wire, items_for_wire = convert_score_media_for_wire(query, items)

# Build request body
request_body: dict[str, Any] = {
Expand Down Expand Up @@ -2337,14 +2343,15 @@ async def extract(
single_item = not isinstance(items, list)
items_list = [items] if single_item else items

# Convert images and documents to wire format (bytes + format hint)
# Convert media and documents to wire format (bytes + format hint)
items_for_wire = []
for item in items_list:
wire_item: dict[str, Any] = {**item} # ty: ignore[invalid-argument-type]
if "images" in wire_item:
wire_item = convert_item_images(wire_item)
if "document" in wire_item:
wire_item = convert_item_document(wire_item)
wire_item = convert_item_media(wire_item)
items_for_wire.append(wire_item)

# Build request body
Expand Down
25 changes: 16 additions & 9 deletions packages/sie_sdk/src/sie_sdk/client/sync.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,7 @@
from sie_sdk.files import resolve_upload
from sie_sdk.images import convert_item_images
from sie_sdk.jobs import TERMINAL_JOB_STATES, build_job_body, decode_chunk_bytes, job_chunks
from sie_sdk.media import convert_item_media
from sie_sdk.types import (
Batch,
CapacityInfo,
Expand Down Expand Up @@ -104,7 +105,7 @@
check_version_skew,
compute_oom_backoff,
compute_retry_delay,
convert_score_images_for_wire,
convert_score_media_for_wire,
get_error_code,
get_retry_after,
get_sdk_version,
Expand Down Expand Up @@ -1081,12 +1082,17 @@ def encode(
single_item = not isinstance(items, list)
items_list = [items] if single_item else items

# Convert images to JPEG bytes for transport.
# Only copy items that have images — text-only items are passed through directly
items_for_wire = [
convert_item_images({**item}) if "images" in item else item # ty: ignore[invalid-argument-type]
for item in items_list
]
# Convert media to bytes for transport.
# Only copy media-bearing items; text-only items pass through directly.
items_for_wire = []
for item in items_list:
if not any(field in item for field in ("images", "audio", "video")):
items_for_wire.append(item)
continue
wire_item: dict[str, Any] = {**item} # ty: ignore[invalid-argument-type]
if "images" in wire_item:
wire_item = convert_item_images(wire_item)
items_for_wire.append(convert_item_media(wire_item))

# Build request body
request_body: dict[str, Any] = {"items": items_for_wire}
Expand Down Expand Up @@ -1729,7 +1735,7 @@ def score(
# Resolve defaults and pool
pool_name, resolved_gpu = self._resolve_pool_and_gpu(gpu)
resolved_options = self._resolve_options(options)
query_for_wire, items_for_wire = convert_score_images_for_wire(query, items)
query_for_wire, items_for_wire = convert_score_media_for_wire(query, items)

# Build request body
request_body: dict[str, Any] = {
Expand Down Expand Up @@ -2569,14 +2575,15 @@ def extract(
single_item = not isinstance(items, list)
items_list = [items] if single_item else items

# Convert images and documents to wire format (bytes + format hint)
# Convert media and documents to wire format (bytes + format hint)
items_for_wire = []
for item in items_list:
wire_item: dict[str, Any] = {**item} # ty: ignore[invalid-argument-type]
if "images" in wire_item:
wire_item = convert_item_images(wire_item)
if "document" in wire_item:
wire_item = convert_item_document(wire_item)
wire_item = convert_item_media(wire_item)
items_for_wire.append(wire_item)

# Build request body
Expand Down
Loading