Skip to content
Merged
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
74 changes: 64 additions & 10 deletions extensions/business/edge_inference_api/base_inference_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -2596,6 +2596,39 @@ def maybe_mark_request_timeout(self, request_id: str, request_data: Dict[str, An
self._decrement_active_requests()
return True

def _extract_request_id_override(self, kwargs: Dict[str, Any]):
"""Extract an optional caller-provided request id from endpoint kwargs.

Existing callers do not pass a request id and keep the generated-id path.
New paired apps may pass `request_id` so wrapper and inference tracking use
the same id without a separate mapping store.
"""
if not isinstance(kwargs, dict):
return None
values = []
for key in ('request_id', 'REQUEST_ID'):
if key in kwargs:
value = kwargs.pop(key)
if value is not None:
values.append(value)
if not values:
return None
request_id = values[0]
for value in values[1:]:
if value != request_id:
raise ValueError("Conflicting request_id and REQUEST_ID values.")
if not isinstance(request_id, str):
raise ValueError("request_id must be a string.")
request_id = request_id.strip()
if not request_id:
raise ValueError("request_id must not be empty.")
if len(request_id) > 256:
raise ValueError("request_id must not exceed 256 characters.")
allowed_chars = set("abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789._:-")
if any(ch not in allowed_chars for ch in request_id):
raise ValueError("request_id contains unsupported characters.")
return request_id

def solve_postponed_request(self, request_id: str):
"""
Resolve or requeue a postponed request by checking its current status.
Expand Down Expand Up @@ -2666,6 +2699,8 @@ def register_request(
Generated request_id and the stored request data dictionary.
"""
request_id = request_id or self.uuid()
if request_id in self._requests:
raise ValueError(f"Request ID {request_id} already exists.")
start_time = self.time()
request_data = {
"request_id": request_id,
Expand Down Expand Up @@ -2857,6 +2892,7 @@ def predict(
def predict_async(
self,
authorization: Optional[str] = None,
request_id: Optional[str] = None,
**kwargs
):
"""
Expand All @@ -2866,6 +2902,9 @@ def predict_async(
----------
authorization : str or None, optional
Authorization token supplied by the caller.
request_id : str or None, optional
Caller-provided id to use for request tracking. If omitted, the API
keeps the legacy generated-id behavior.
**kwargs
Additional parameters forwarded to request handling.

Expand All @@ -2877,6 +2916,7 @@ def predict_async(
return self._predict_entrypoint(
authorization=authorization,
async_request=True,
request_id=request_id,
**kwargs
)
"""END API ENDPOINTS"""
Expand Down Expand Up @@ -2992,6 +3032,20 @@ def _predict_entrypoint(
return {'error': f"Unexpected error: {str(exc)}", 'status': 'error'}
# endtry

request_id_override = None
if delegated_execution:
request_id_override = delegation_context.get('delegation_id')
elif async_request:
try:
request_id_override = self._extract_request_id_override(kwargs)
except ValueError as exc:
return {'error': str(exc), 'status': 'error'}
else:
# `request_id` is reserved for async tracking and must not leak into
# synchronous model parameters when clients reuse async payload shapes.
kwargs.pop('request_id', None)
kwargs.pop('REQUEST_ID', None)

err = self.check_predict_params(**kwargs)
if err is not None:
return {'error': err}
Expand All @@ -3000,16 +3054,16 @@ def _predict_entrypoint(
if 'metadata' in parameters:
metadata = parameters.pop('metadata') or {}
# endif 'metadata' in parameters
request_id_override = None
if delegated_execution:
request_id_override = delegation_context.get('delegation_id')
request_id, request_data = self.register_request(
subject=subject,
parameters=parameters,
metadata=metadata,
timeout=parameters.get('timeout'),
request_id=request_id_override,
)
try:
request_id, request_data = self.register_request(
subject=subject,
parameters=parameters,
metadata=metadata,
timeout=parameters.get('timeout'),
request_id=request_id_override,
)
except ValueError as exc:
return {'error': str(exc), 'status': 'error'}
request_data['endpoint_name'] = endpoint_name
request_data['async_request'] = async_request

Expand Down
5 changes: 5 additions & 0 deletions extensions/business/edge_inference_api/cv_inference_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -236,6 +236,7 @@ def predict_async(
image_data: str = '',
metadata: Optional[Dict[str, Any]] = None,
authorization: Optional[str] = None,
request_id: Optional[str] = None,
**kwargs
):
"""
Expand All @@ -249,6 +250,9 @@ def predict_async(
Optional metadata accompanying the request.
authorization : str or None, optional
Bearer token used for authentication.
request_id : str or None, optional
Caller-provided id to use for request tracking. If omitted, the API
keeps the legacy generated-id behavior.
**kwargs
Extra parameters forwarded to the base handler.

Expand All @@ -261,6 +265,7 @@ def predict_async(
image_data=image_data,
metadata=metadata,
authorization=authorization,
request_id=request_id,
**kwargs
)
"""END API ENDPOINTS"""
Expand Down
5 changes: 5 additions & 0 deletions extensions/business/edge_inference_api/llm_inference_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -399,6 +399,7 @@ def predict_async(
response_format: Optional[Dict[str, Any]] = None,
metadata: Optional[Dict[str, Any]] = None,
authorization: Optional[str] = None,
request_id: Optional[str] = None,
**kwargs
):
"""
Expand All @@ -422,6 +423,9 @@ def predict_async(
Additional metadata to store with the request.
authorization : str or None, optional
Bearer token used for authentication.
request_id : str or None, optional
Caller-provided id to use for request tracking. If omitted, the API
keeps the legacy generated-id behavior.
**kwargs
Extra parameters forwarded to the base handler.

Expand All @@ -439,6 +443,7 @@ def predict_async(
response_format=response_format,
metadata=metadata,
authorization=authorization,
request_id=request_id,
**kwargs
)

Expand Down
5 changes: 5 additions & 0 deletions extensions/business/edge_inference_api/sd_inference_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -283,6 +283,7 @@ def predict_async(
struct_data: Any = None,
metadata: Optional[Dict[str, Any]] = None,
authorization: Optional[str] = None,
request_id: Optional[str] = None,
**kwargs
):
"""
Expand All @@ -296,6 +297,9 @@ def predict_async(
Optional metadata accompanying the request.
authorization : str or None, optional
Bearer token used for authentication.
request_id : str or None, optional
Caller-provided id to use for request tracking. If omitted, the API
keeps the legacy generated-id behavior.
**kwargs
Extra parameters forwarded to the base handler.

Expand All @@ -308,6 +312,7 @@ def predict_async(
struct_data=struct_data,
metadata=metadata,
authorization=authorization,
request_id=request_id,
**kwargs
)
"""END API ENDPOINTS"""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -580,6 +580,88 @@ def test_predict_entrypoint_queues_when_full_and_no_peer(self):
request_id = result["request_id"]
self.assertEqual(plugin._requests[request_id]["queue_state"], "queued") # pylint: disable=protected-access

def test_predict_entrypoint_accepts_caller_request_id(self):
plugin = self._make_plugin()

result = plugin._predict_entrypoint( # pylint: disable=protected-access
authorization=None,
async_request=True,
request_id="client-req_1.2:3",
metadata={"source": "test"},
)

self.assertEqual(result["request_id"], "client-req_1.2:3")
self.assertIn("client-req_1.2:3", plugin._requests) # pylint: disable=protected-access
self.assertEqual(plugin.payloads[-1]["REQUEST_ID"], "client-req_1.2:3")
self.assertNotIn("request_id", plugin._requests["client-req_1.2:3"]["parameters"]) # pylint: disable=protected-access

def test_predict_async_without_request_id_keeps_generated_id_behavior(self):
plugin = self._make_plugin()

result = plugin.predict_async(authorization=None, request_id=None)

self.assertEqual(result["request_id"], "req-1")
self.assertIn("req-1", plugin._requests) # pylint: disable=protected-access

def test_sync_predict_does_not_treat_request_id_as_tracking_override(self):
plugin = self._make_plugin()

plugin._predict_entrypoint( # pylint: disable=protected-access
authorization=None,
async_request=False,
request_id="client-sync-id",
REQUEST_ID="client-sync-id-upper",
)

self.assertIn("req-1", plugin._requests) # pylint: disable=protected-access
self.assertNotIn("client-sync-id", plugin._requests) # pylint: disable=protected-access
self.assertNotIn("request_id", plugin._requests["req-1"]["parameters"]) # pylint: disable=protected-access
self.assertNotIn("REQUEST_ID", plugin._requests["req-1"]["parameters"]) # pylint: disable=protected-access
self.assertEqual(plugin.payloads[-1]["REQUEST_ID"], "req-1")

def test_predict_entrypoint_accepts_uppercase_request_id_alias(self):
plugin = self._make_plugin()

result = plugin._predict_entrypoint( # pylint: disable=protected-access
authorization=None,
async_request=True,
REQUEST_ID="client-req-2",
)

self.assertEqual(result["request_id"], "client-req-2")
self.assertIn("client-req-2", plugin._requests) # pylint: disable=protected-access

def test_predict_entrypoint_rejects_duplicate_caller_request_id(self):
plugin = self._make_plugin()

first = plugin._predict_entrypoint( # pylint: disable=protected-access
authorization=None,
async_request=True,
request_id="client-req-dup",
)
second = plugin._predict_entrypoint( # pylint: disable=protected-access
authorization=None,
async_request=True,
request_id="client-req-dup",
)

self.assertEqual(first["request_id"], "client-req-dup")
self.assertEqual(second["status"], "error")
self.assertIn("already exists", second["error"])

def test_predict_entrypoint_rejects_invalid_caller_request_id(self):
plugin = self._make_plugin()

result = plugin._predict_entrypoint( # pylint: disable=protected-access
authorization=None,
async_request=True,
request_id="../bad",
)

self.assertEqual(result["status"], "error")
self.assertIn("unsupported characters", result["error"])
self.assertNotIn("../bad", plugin._requests) # pylint: disable=protected-access

def test_predict_entrypoint_fails_cleanly_when_delegated_request_cannot_encode(self):
plugin = self._make_plugin()
plugin._active_execution_slots.add("busy") # pylint: disable=protected-access
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -276,6 +276,7 @@ def predict_async(
text: str = "",
metadata: Optional[Dict[str, Any]] = None,
authorization: Optional[str] = None,
request_id: Optional[str] = None,
**kwargs
):
"""
Expand All @@ -289,6 +290,9 @@ def predict_async(
Optional metadata accompanying the request.
authorization : str or None, optional
Bearer token used for authentication.
request_id : str or None, optional
Caller-provided id to use for request tracking. If omitted, the API
keeps the legacy generated-id behavior.
**kwargs
Extra parameters forwarded to the base handler.

Expand All @@ -301,6 +305,7 @@ def predict_async(
text=text,
metadata=metadata,
authorization=authorization,
request_id=request_id,
**kwargs
)
"""END API ENDPOINTS"""
Expand Down
2 changes: 1 addition & 1 deletion ver.py
Original file line number Diff line number Diff line change
@@ -1 +1 @@
__VER__ = '2.10.290'
__VER__ = '2.10.300'
Loading