diff --git a/extensions/business/edge_inference_api/base_inference_api.py b/extensions/business/edge_inference_api/base_inference_api.py index bef7290c..4ba4003c 100644 --- a/extensions/business/edge_inference_api/base_inference_api.py +++ b/extensions/business/edge_inference_api/base_inference_api.py @@ -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. @@ -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, @@ -2857,6 +2892,7 @@ def predict( def predict_async( self, authorization: Optional[str] = None, + request_id: Optional[str] = None, **kwargs ): """ @@ -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. @@ -2877,6 +2916,7 @@ def predict_async( return self._predict_entrypoint( authorization=authorization, async_request=True, + request_id=request_id, **kwargs ) """END API ENDPOINTS""" @@ -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} @@ -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 diff --git a/extensions/business/edge_inference_api/cv_inference_api.py b/extensions/business/edge_inference_api/cv_inference_api.py index 365b1336..5b53b9a3 100644 --- a/extensions/business/edge_inference_api/cv_inference_api.py +++ b/extensions/business/edge_inference_api/cv_inference_api.py @@ -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 ): """ @@ -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. @@ -261,6 +265,7 @@ def predict_async( image_data=image_data, metadata=metadata, authorization=authorization, + request_id=request_id, **kwargs ) """END API ENDPOINTS""" diff --git a/extensions/business/edge_inference_api/llm_inference_api.py b/extensions/business/edge_inference_api/llm_inference_api.py index 429cb319..2b4f5019 100644 --- a/extensions/business/edge_inference_api/llm_inference_api.py +++ b/extensions/business/edge_inference_api/llm_inference_api.py @@ -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 ): """ @@ -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. @@ -439,6 +443,7 @@ def predict_async( response_format=response_format, metadata=metadata, authorization=authorization, + request_id=request_id, **kwargs ) diff --git a/extensions/business/edge_inference_api/sd_inference_api.py b/extensions/business/edge_inference_api/sd_inference_api.py index 13c39585..f05310ee 100644 --- a/extensions/business/edge_inference_api/sd_inference_api.py +++ b/extensions/business/edge_inference_api/sd_inference_api.py @@ -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 ): """ @@ -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. @@ -308,6 +312,7 @@ def predict_async( struct_data=struct_data, metadata=metadata, authorization=authorization, + request_id=request_id, **kwargs ) """END API ENDPOINTS""" diff --git a/extensions/business/edge_inference_api/test_base_inference_api_balancing.py b/extensions/business/edge_inference_api/test_base_inference_api_balancing.py index e0662a1c..e9c9c0d7 100644 --- a/extensions/business/edge_inference_api/test_base_inference_api_balancing.py +++ b/extensions/business/edge_inference_api/test_base_inference_api_balancing.py @@ -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 diff --git a/extensions/business/edge_inference_api/text_classifier_inference_api.py b/extensions/business/edge_inference_api/text_classifier_inference_api.py index 106b31b8..0db5f827 100644 --- a/extensions/business/edge_inference_api/text_classifier_inference_api.py +++ b/extensions/business/edge_inference_api/text_classifier_inference_api.py @@ -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 ): """ @@ -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. @@ -301,6 +305,7 @@ def predict_async( text=text, metadata=metadata, authorization=authorization, + request_id=request_id, **kwargs ) """END API ENDPOINTS""" diff --git a/ver.py b/ver.py index 31580534..9f49e4d0 100644 --- a/ver.py +++ b/ver.py @@ -1 +1 @@ -__VER__ = '2.10.290' +__VER__ = '2.10.300'