diff --git a/framework/proto/flwr/proto/runtime.proto b/framework/proto/flwr/proto/runtime.proto index 040c7b0ff2f7..5bbf029166b1 100644 --- a/framework/proto/flwr/proto/runtime.proto +++ b/framework/proto/flwr/proto/runtime.proto @@ -152,7 +152,11 @@ message PushTaskEventsRequest { repeated TaskEvent events = 1; } message PushTaskEventsResponse {} // PullTaskMessage messages -message PullTaskMessageRequest { optional uint64 limit = 1; } +message PullTaskMessageRequest { + optional uint64 limit = 1; + // Filter concurrent exchanges by sender task ID. + optional uint64 src_task_id = 2; +} message PullTaskMessageResponse { repeated Message messages = 1; } // RecordTaskUsage messages diff --git a/framework/py/flwr/proto/runtime_pb2.py b/framework/py/flwr/proto/runtime_pb2.py index fd54800c1811..81f3fb2ad343 100644 --- a/framework/py/flwr/proto/runtime_pb2.py +++ b/framework/py/flwr/proto/runtime_pb2.py @@ -32,7 +32,7 @@ from flwr.proto import task_pb2 as flwr_dot_proto_dot_task__pb2 -DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x18\x66lwr/proto/runtime.proto\x12\nflwr.proto\x1a\x18\x66lwr/proto/control.proto\x1a\x14\x66lwr/proto/fab.proto\x1a\"flwr/proto/federation_config.proto\x1a\x14\x66lwr/proto/log.proto\x1a\x18\x66lwr/proto/message.proto\x1a\x15\x66lwr/proto/node.proto\x1a\x14\x66lwr/proto/run.proto\x1a\x15\x66lwr/proto/task.proto\"\x19\n\x17PullPendingTasksRequest\";\n\x18PullPendingTasksResponse\x12\x1f\n\x05tasks\x18\x01 \x03(\x0b\x32\x10.flwr.proto.Task\"#\n\x10\x43laimTaskRequest\x12\x0f\n\x07task_id\x18\x01 \x01(\x04\"1\n\x11\x43laimTaskResponse\x12\x12\n\x05token\x18\x01 \x01(\tH\x00\x88\x01\x01\x42\x08\n\x06_token\"\x1a\n\x18SendTaskHeartbeatRequest\",\n\x19SendTaskHeartbeatResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\"z\n\x16PushAppMessagesRequest\x12*\n\rmessages_list\x18\x02 \x03(\x0b\x32\x13.flwr.proto.Message\x12\x34\n\x14message_object_trees\x18\x03 \x03(\x0b\x32\x16.flwr.proto.ObjectTree\"[\n\x17PushAppMessagesResponse\x12\x13\n\x0bmessage_ids\x18\x01 \x03(\t\x12\x17\n\x0fobjects_to_push\x18\x02 \x03(\t\x12\x12\n\nsession_id\x18\x03 \x01(\t\"-\n\x16PullAppMessagesRequest\x12\x13\n\x0bmessage_ids\x18\x02 \x03(\t\"{\n\x17PullAppMessagesResponse\x12*\n\rmessages_list\x18\x01 \x03(\x0b\x32\x13.flwr.proto.Message\x12\x34\n\x14message_object_trees\x18\x02 \x03(\x0b\x32\x16.flwr.proto.ObjectTree\"\x11\n\x0fGetNodesRequest\"3\n\x10GetNodesResponse\x12\x1f\n\x05nodes\x18\x01 \x03(\x0b\x32\x10.flwr.proto.Node\">\n\x16PushTaskMessageRequest\x12$\n\x07message\x18\x01 \x01(\x0b\x32\x13.flwr.proto.Message\"-\n\x17PushTaskMessageResponse\x12\x12\n\nmessage_id\x18\x01 \x01(\t\">\n\x15PushTaskEventsRequest\x12%\n\x06\x65vents\x18\x01 \x03(\x0b\x32\x15.flwr.proto.TaskEvent\"\x18\n\x16PushTaskEventsResponse\"6\n\x16PullTaskMessageRequest\x12\x12\n\x05limit\x18\x01 \x01(\x04H\x00\x88\x01\x01\x42\x08\n\x06_limit\"@\n\x17PullTaskMessageResponse\x12%\n\x08messages\x18\x01 \x03(\x0b\x32\x13.flwr.proto.Message\"C\n\x16RecordTaskUsageRequest\x12)\n\ntask_usage\x18\x01 \x01(\x0b\x32\x15.flwr.proto.TaskUsage\"\x19\n\x17RecordTaskUsageResponse\"\x15\n\x13GetConnectorRequest\"\\\n\x14GetConnectorResponse\x12\x15\n\rconnector_ref\x18\x01 \x01(\t\x12\x18\n\x10\x63redentials_json\x18\x02 \x01(\t\x12\x13\n\x0b\x63onfig_json\x18\x03 \x01(\t\"\x16\n\x14PullTaskInputRequest\"\xc3\x01\n\x15PullTaskInputResponse\x12$\n\x07\x63ontext\x18\x01 \x01(\x0b\x32\x13.flwr.proto.Context\x12\x1c\n\x03run\x18\x02 \x01(\x0b\x32\x0f.flwr.proto.Run\x12\x1c\n\x03\x66\x61\x62\x18\x03 \x01(\x0b\x32\x0f.flwr.proto.Fab\x12\x37\n\x11\x66\x65\x64\x65ration_config\x18\x04 \x01(\x0b\x32\x1c.flwr.proto.SimulationConfig\x12\x0f\n\x07task_id\x18\x05 \x01(\x04\"\x98\x01\n\x15PushTaskOutputRequest\x12$\n\x07\x63ontext\x18\x01 \x01(\x0b\x32\x13.flwr.proto.Context\x12\x12\n\nsub_status\x18\x02 \x01(\t\x12\x0f\n\x07\x64\x65tails\x18\x03 \x01(\t\x12\x1e\n\x11\x63lientapp_runtime\x18\x04 \x01(\x01H\x00\x88\x01\x01\x42\x14\n\x12_clientapp_runtime\"\x18\n\x16PushTaskOutputResponse\"\x99\x01\n\x11\x43reateTaskRequest\x12\x0c\n\x04type\x18\x01 \x01(\t\x12\x15\n\x08\x66\x61\x62_hash\x18\x02 \x01(\tH\x00\x88\x01\x01\x12\x16\n\tmodel_ref\x18\x03 \x01(\tH\x01\x88\x01\x01\x12\x1a\n\rconnector_ref\x18\x04 \x01(\tH\x02\x88\x01\x01\x42\x0b\n\t_fab_hashB\x0c\n\n_model_refB\x10\n\x0e_connector_ref\"6\n\x12\x43reateTaskResponse\x12\x14\n\x07task_id\x18\x01 \x01(\x04H\x00\x88\x01\x01\x42\n\n\x08_task_id2\x9d\r\n\x07Runtime\x12_\n\x10PullPendingTasks\x12#.flwr.proto.PullPendingTasksRequest\x1a$.flwr.proto.PullPendingTasksResponse\"\x00\x12J\n\tClaimTask\x12\x1c.flwr.proto.ClaimTaskRequest\x1a\x1d.flwr.proto.ClaimTaskResponse\"\x00\x12\x62\n\x11SendTaskHeartbeat\x12$.flwr.proto.SendTaskHeartbeatRequest\x1a%.flwr.proto.SendTaskHeartbeatResponse\"\x00\x12V\n\rPullTaskInput\x12 .flwr.proto.PullTaskInputRequest\x1a!.flwr.proto.PullTaskInputResponse\"\x00\x12Y\n\x0ePushTaskOutput\x12!.flwr.proto.PushTaskOutputRequest\x1a\".flwr.proto.PushTaskOutputResponse\"\x00\x12M\n\nPushObject\x12\x1d.flwr.proto.PushObjectRequest\x1a\x1e.flwr.proto.PushObjectResponse\"\x00\x12M\n\nPullObject\x12\x1d.flwr.proto.PullObjectRequest\x1a\x1e.flwr.proto.PullObjectResponse\"\x00\x12q\n\x16\x43onfirmMessageReceived\x12).flwr.proto.ConfirmMessageReceivedRequest\x1a*.flwr.proto.ConfirmMessageReceivedResponse\"\x00\x12M\n\nCreateTask\x12\x1d.flwr.proto.CreateTaskRequest\x1a\x1e.flwr.proto.CreateTaskResponse\"\x00\x12\\\n\x0fStartAutomation\x12\".flwr.proto.StartAutomationRequest\x1a#.flwr.proto.StartAutomationResponse\"\x00\x12\\\n\x0fPushTaskMessage\x12\".flwr.proto.PushTaskMessageRequest\x1a#.flwr.proto.PushTaskMessageResponse\"\x00\x12Y\n\x0ePushTaskEvents\x12!.flwr.proto.PushTaskEventsRequest\x1a\".flwr.proto.PushTaskEventsResponse\"\x00\x12\\\n\x0fPullTaskMessage\x12\".flwr.proto.PullTaskMessageRequest\x1a#.flwr.proto.PullTaskMessageResponse\"\x00\x12\\\n\x0fRecordTaskUsage\x12\".flwr.proto.RecordTaskUsageRequest\x1a#.flwr.proto.RecordTaskUsageResponse\"\x00\x12S\n\x0cGetConnector\x12\x1f.flwr.proto.GetConnectorRequest\x1a .flwr.proto.GetConnectorResponse\"\x00\x12G\n\x08PushLogs\x12\x1b.flwr.proto.PushLogsRequest\x1a\x1c.flwr.proto.PushLogsResponse\"\x00\x12Y\n\x0cPushMessages\x12\".flwr.proto.PushAppMessagesRequest\x1a#.flwr.proto.PushAppMessagesResponse\"\x00\x12Y\n\x0cPullMessages\x12\".flwr.proto.PullAppMessagesRequest\x1a#.flwr.proto.PullAppMessagesResponse\"\x00\x12G\n\x08GetNodes\x12\x1b.flwr.proto.GetNodesRequest\x1a\x1c.flwr.proto.GetNodesResponse\"\x00\x62\x06proto3') +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x18\x66lwr/proto/runtime.proto\x12\nflwr.proto\x1a\x18\x66lwr/proto/control.proto\x1a\x14\x66lwr/proto/fab.proto\x1a\"flwr/proto/federation_config.proto\x1a\x14\x66lwr/proto/log.proto\x1a\x18\x66lwr/proto/message.proto\x1a\x15\x66lwr/proto/node.proto\x1a\x14\x66lwr/proto/run.proto\x1a\x15\x66lwr/proto/task.proto\"\x19\n\x17PullPendingTasksRequest\";\n\x18PullPendingTasksResponse\x12\x1f\n\x05tasks\x18\x01 \x03(\x0b\x32\x10.flwr.proto.Task\"#\n\x10\x43laimTaskRequest\x12\x0f\n\x07task_id\x18\x01 \x01(\x04\"1\n\x11\x43laimTaskResponse\x12\x12\n\x05token\x18\x01 \x01(\tH\x00\x88\x01\x01\x42\x08\n\x06_token\"\x1a\n\x18SendTaskHeartbeatRequest\",\n\x19SendTaskHeartbeatResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\"z\n\x16PushAppMessagesRequest\x12*\n\rmessages_list\x18\x02 \x03(\x0b\x32\x13.flwr.proto.Message\x12\x34\n\x14message_object_trees\x18\x03 \x03(\x0b\x32\x16.flwr.proto.ObjectTree\"[\n\x17PushAppMessagesResponse\x12\x13\n\x0bmessage_ids\x18\x01 \x03(\t\x12\x17\n\x0fobjects_to_push\x18\x02 \x03(\t\x12\x12\n\nsession_id\x18\x03 \x01(\t\"-\n\x16PullAppMessagesRequest\x12\x13\n\x0bmessage_ids\x18\x02 \x03(\t\"{\n\x17PullAppMessagesResponse\x12*\n\rmessages_list\x18\x01 \x03(\x0b\x32\x13.flwr.proto.Message\x12\x34\n\x14message_object_trees\x18\x02 \x03(\x0b\x32\x16.flwr.proto.ObjectTree\"\x11\n\x0fGetNodesRequest\"3\n\x10GetNodesResponse\x12\x1f\n\x05nodes\x18\x01 \x03(\x0b\x32\x10.flwr.proto.Node\">\n\x16PushTaskMessageRequest\x12$\n\x07message\x18\x01 \x01(\x0b\x32\x13.flwr.proto.Message\"-\n\x17PushTaskMessageResponse\x12\x12\n\nmessage_id\x18\x01 \x01(\t\">\n\x15PushTaskEventsRequest\x12%\n\x06\x65vents\x18\x01 \x03(\x0b\x32\x15.flwr.proto.TaskEvent\"\x18\n\x16PushTaskEventsResponse\"`\n\x16PullTaskMessageRequest\x12\x12\n\x05limit\x18\x01 \x01(\x04H\x00\x88\x01\x01\x12\x18\n\x0bsrc_task_id\x18\x02 \x01(\x04H\x01\x88\x01\x01\x42\x08\n\x06_limitB\x0e\n\x0c_src_task_id\"@\n\x17PullTaskMessageResponse\x12%\n\x08messages\x18\x01 \x03(\x0b\x32\x13.flwr.proto.Message\"C\n\x16RecordTaskUsageRequest\x12)\n\ntask_usage\x18\x01 \x01(\x0b\x32\x15.flwr.proto.TaskUsage\"\x19\n\x17RecordTaskUsageResponse\"\x15\n\x13GetConnectorRequest\"\\\n\x14GetConnectorResponse\x12\x15\n\rconnector_ref\x18\x01 \x01(\t\x12\x18\n\x10\x63redentials_json\x18\x02 \x01(\t\x12\x13\n\x0b\x63onfig_json\x18\x03 \x01(\t\"\x16\n\x14PullTaskInputRequest\"\xc3\x01\n\x15PullTaskInputResponse\x12$\n\x07\x63ontext\x18\x01 \x01(\x0b\x32\x13.flwr.proto.Context\x12\x1c\n\x03run\x18\x02 \x01(\x0b\x32\x0f.flwr.proto.Run\x12\x1c\n\x03\x66\x61\x62\x18\x03 \x01(\x0b\x32\x0f.flwr.proto.Fab\x12\x37\n\x11\x66\x65\x64\x65ration_config\x18\x04 \x01(\x0b\x32\x1c.flwr.proto.SimulationConfig\x12\x0f\n\x07task_id\x18\x05 \x01(\x04\"\x98\x01\n\x15PushTaskOutputRequest\x12$\n\x07\x63ontext\x18\x01 \x01(\x0b\x32\x13.flwr.proto.Context\x12\x12\n\nsub_status\x18\x02 \x01(\t\x12\x0f\n\x07\x64\x65tails\x18\x03 \x01(\t\x12\x1e\n\x11\x63lientapp_runtime\x18\x04 \x01(\x01H\x00\x88\x01\x01\x42\x14\n\x12_clientapp_runtime\"\x18\n\x16PushTaskOutputResponse\"\x99\x01\n\x11\x43reateTaskRequest\x12\x0c\n\x04type\x18\x01 \x01(\t\x12\x15\n\x08\x66\x61\x62_hash\x18\x02 \x01(\tH\x00\x88\x01\x01\x12\x16\n\tmodel_ref\x18\x03 \x01(\tH\x01\x88\x01\x01\x12\x1a\n\rconnector_ref\x18\x04 \x01(\tH\x02\x88\x01\x01\x42\x0b\n\t_fab_hashB\x0c\n\n_model_refB\x10\n\x0e_connector_ref\"6\n\x12\x43reateTaskResponse\x12\x14\n\x07task_id\x18\x01 \x01(\x04H\x00\x88\x01\x01\x42\n\n\x08_task_id2\x9d\r\n\x07Runtime\x12_\n\x10PullPendingTasks\x12#.flwr.proto.PullPendingTasksRequest\x1a$.flwr.proto.PullPendingTasksResponse\"\x00\x12J\n\tClaimTask\x12\x1c.flwr.proto.ClaimTaskRequest\x1a\x1d.flwr.proto.ClaimTaskResponse\"\x00\x12\x62\n\x11SendTaskHeartbeat\x12$.flwr.proto.SendTaskHeartbeatRequest\x1a%.flwr.proto.SendTaskHeartbeatResponse\"\x00\x12V\n\rPullTaskInput\x12 .flwr.proto.PullTaskInputRequest\x1a!.flwr.proto.PullTaskInputResponse\"\x00\x12Y\n\x0ePushTaskOutput\x12!.flwr.proto.PushTaskOutputRequest\x1a\".flwr.proto.PushTaskOutputResponse\"\x00\x12M\n\nPushObject\x12\x1d.flwr.proto.PushObjectRequest\x1a\x1e.flwr.proto.PushObjectResponse\"\x00\x12M\n\nPullObject\x12\x1d.flwr.proto.PullObjectRequest\x1a\x1e.flwr.proto.PullObjectResponse\"\x00\x12q\n\x16\x43onfirmMessageReceived\x12).flwr.proto.ConfirmMessageReceivedRequest\x1a*.flwr.proto.ConfirmMessageReceivedResponse\"\x00\x12M\n\nCreateTask\x12\x1d.flwr.proto.CreateTaskRequest\x1a\x1e.flwr.proto.CreateTaskResponse\"\x00\x12\\\n\x0fStartAutomation\x12\".flwr.proto.StartAutomationRequest\x1a#.flwr.proto.StartAutomationResponse\"\x00\x12\\\n\x0fPushTaskMessage\x12\".flwr.proto.PushTaskMessageRequest\x1a#.flwr.proto.PushTaskMessageResponse\"\x00\x12Y\n\x0ePushTaskEvents\x12!.flwr.proto.PushTaskEventsRequest\x1a\".flwr.proto.PushTaskEventsResponse\"\x00\x12\\\n\x0fPullTaskMessage\x12\".flwr.proto.PullTaskMessageRequest\x1a#.flwr.proto.PullTaskMessageResponse\"\x00\x12\\\n\x0fRecordTaskUsage\x12\".flwr.proto.RecordTaskUsageRequest\x1a#.flwr.proto.RecordTaskUsageResponse\"\x00\x12S\n\x0cGetConnector\x12\x1f.flwr.proto.GetConnectorRequest\x1a .flwr.proto.GetConnectorResponse\"\x00\x12G\n\x08PushLogs\x12\x1b.flwr.proto.PushLogsRequest\x1a\x1c.flwr.proto.PushLogsResponse\"\x00\x12Y\n\x0cPushMessages\x12\".flwr.proto.PushAppMessagesRequest\x1a#.flwr.proto.PushAppMessagesResponse\"\x00\x12Y\n\x0cPullMessages\x12\".flwr.proto.PullAppMessagesRequest\x1a#.flwr.proto.PullAppMessagesResponse\"\x00\x12G\n\x08GetNodes\x12\x1b.flwr.proto.GetNodesRequest\x1a\x1c.flwr.proto.GetNodesResponse\"\x00\x62\x06proto3') _globals = globals() _builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) @@ -72,29 +72,29 @@ _globals['_PUSHTASKEVENTSRESPONSE']._serialized_start=1126 _globals['_PUSHTASKEVENTSRESPONSE']._serialized_end=1150 _globals['_PULLTASKMESSAGEREQUEST']._serialized_start=1152 - _globals['_PULLTASKMESSAGEREQUEST']._serialized_end=1206 - _globals['_PULLTASKMESSAGERESPONSE']._serialized_start=1208 - _globals['_PULLTASKMESSAGERESPONSE']._serialized_end=1272 - _globals['_RECORDTASKUSAGEREQUEST']._serialized_start=1274 - _globals['_RECORDTASKUSAGEREQUEST']._serialized_end=1341 - _globals['_RECORDTASKUSAGERESPONSE']._serialized_start=1343 - _globals['_RECORDTASKUSAGERESPONSE']._serialized_end=1368 - _globals['_GETCONNECTORREQUEST']._serialized_start=1370 - _globals['_GETCONNECTORREQUEST']._serialized_end=1391 - _globals['_GETCONNECTORRESPONSE']._serialized_start=1393 - _globals['_GETCONNECTORRESPONSE']._serialized_end=1485 - _globals['_PULLTASKINPUTREQUEST']._serialized_start=1487 - _globals['_PULLTASKINPUTREQUEST']._serialized_end=1509 - _globals['_PULLTASKINPUTRESPONSE']._serialized_start=1512 - _globals['_PULLTASKINPUTRESPONSE']._serialized_end=1707 - _globals['_PUSHTASKOUTPUTREQUEST']._serialized_start=1710 - _globals['_PUSHTASKOUTPUTREQUEST']._serialized_end=1862 - _globals['_PUSHTASKOUTPUTRESPONSE']._serialized_start=1864 - _globals['_PUSHTASKOUTPUTRESPONSE']._serialized_end=1888 - _globals['_CREATETASKREQUEST']._serialized_start=1891 - _globals['_CREATETASKREQUEST']._serialized_end=2044 - _globals['_CREATETASKRESPONSE']._serialized_start=2046 - _globals['_CREATETASKRESPONSE']._serialized_end=2100 - _globals['_RUNTIME']._serialized_start=2103 - _globals['_RUNTIME']._serialized_end=3796 + _globals['_PULLTASKMESSAGEREQUEST']._serialized_end=1248 + _globals['_PULLTASKMESSAGERESPONSE']._serialized_start=1250 + _globals['_PULLTASKMESSAGERESPONSE']._serialized_end=1314 + _globals['_RECORDTASKUSAGEREQUEST']._serialized_start=1316 + _globals['_RECORDTASKUSAGEREQUEST']._serialized_end=1383 + _globals['_RECORDTASKUSAGERESPONSE']._serialized_start=1385 + _globals['_RECORDTASKUSAGERESPONSE']._serialized_end=1410 + _globals['_GETCONNECTORREQUEST']._serialized_start=1412 + _globals['_GETCONNECTORREQUEST']._serialized_end=1433 + _globals['_GETCONNECTORRESPONSE']._serialized_start=1435 + _globals['_GETCONNECTORRESPONSE']._serialized_end=1527 + _globals['_PULLTASKINPUTREQUEST']._serialized_start=1529 + _globals['_PULLTASKINPUTREQUEST']._serialized_end=1551 + _globals['_PULLTASKINPUTRESPONSE']._serialized_start=1554 + _globals['_PULLTASKINPUTRESPONSE']._serialized_end=1749 + _globals['_PUSHTASKOUTPUTREQUEST']._serialized_start=1752 + _globals['_PUSHTASKOUTPUTREQUEST']._serialized_end=1904 + _globals['_PUSHTASKOUTPUTRESPONSE']._serialized_start=1906 + _globals['_PUSHTASKOUTPUTRESPONSE']._serialized_end=1930 + _globals['_CREATETASKREQUEST']._serialized_start=1933 + _globals['_CREATETASKREQUEST']._serialized_end=2086 + _globals['_CREATETASKRESPONSE']._serialized_start=2088 + _globals['_CREATETASKRESPONSE']._serialized_end=2142 + _globals['_RUNTIME']._serialized_start=2145 + _globals['_RUNTIME']._serialized_end=3838 # @@protoc_insertion_point(module_scope) diff --git a/framework/py/flwr/proto/runtime_pb2.pyi b/framework/py/flwr/proto/runtime_pb2.pyi index db6a42b4f8fa..d31f4746129b 100644 --- a/framework/py/flwr/proto/runtime_pb2.pyi +++ b/framework/py/flwr/proto/runtime_pb2.pyi @@ -301,15 +301,22 @@ class PullTaskMessageRequest(google.protobuf.message.Message): DESCRIPTOR: google.protobuf.descriptor.Descriptor LIMIT_FIELD_NUMBER: builtins.int + SRC_TASK_ID_FIELD_NUMBER: builtins.int limit: builtins.int + src_task_id: builtins.int + """Filter concurrent exchanges by sender task ID.""" def __init__( self, *, limit: builtins.int | None = ..., + src_task_id: builtins.int | None = ..., ) -> None: ... - def HasField(self, field_name: typing.Literal["_limit", b"_limit", "limit", b"limit"]) -> builtins.bool: ... - def ClearField(self, field_name: typing.Literal["_limit", b"_limit", "limit", b"limit"]) -> None: ... + def HasField(self, field_name: typing.Literal["_limit", b"_limit", "_src_task_id", b"_src_task_id", "limit", b"limit", "src_task_id", b"src_task_id"]) -> builtins.bool: ... + def ClearField(self, field_name: typing.Literal["_limit", b"_limit", "_src_task_id", b"_src_task_id", "limit", b"limit", "src_task_id", b"src_task_id"]) -> None: ... + @typing.overload def WhichOneof(self, oneof_group: typing.Literal["_limit", b"_limit"]) -> typing.Literal["limit"] | None: ... + @typing.overload + def WhichOneof(self, oneof_group: typing.Literal["_src_task_id", b"_src_task_id"]) -> typing.Literal["src_task_id"] | None: ... global___PullTaskMessageRequest = PullTaskMessageRequest diff --git a/framework/py/flwr/supercore/corestate/corestate.py b/framework/py/flwr/supercore/corestate/corestate.py index 9d645cdbfb62..c69ae6fcfb59 100644 --- a/framework/py/flwr/supercore/corestate/corestate.py +++ b/framework/py/flwr/supercore/corestate/corestate.py @@ -879,6 +879,7 @@ def get_task_message( self, *, dst_task_ids: Sequence[int] | None = None, + src_task_ids: Sequence[int] | None = None, limit: int | None = None, order_by: Literal["created_at"] | None = None, ) -> Sequence[Message]: @@ -891,6 +892,8 @@ def get_task_message( ---------- dst_task_ids : Optional[Sequence[int]] (default: None) Sequence of destination task IDs to filter by. + src_task_ids : Optional[Sequence[int]] (default: None) + Sequence of source task IDs to filter by. limit : Optional[int] (default: None) Maximum number of messages to return. If `None`, no limit is applied. order_by : Optional[Literal["created_at"]] (default: None) @@ -929,6 +932,7 @@ def get_task_events( self, *, run_id: int | None = None, + task_ids: Sequence[int] | None = None, after_task_event_id: int | None = None, ) -> Sequence[TaskEvent]: """Return task-produced run events matching the filters. @@ -938,6 +942,8 @@ def get_task_events( run_id : Optional[int] (default: None) If set, return only events for this run. If set to `None`, return events for all runs. + task_ids : Optional[Sequence[int]] (default: None) + If set, return only events produced by these tasks. after_task_event_id : Optional[int] (default: None) Return only events with an ID greater than this cursor. If set to `None`, retrieve all events. diff --git a/framework/py/flwr/supercore/corestate/corestate_test.py b/framework/py/flwr/supercore/corestate/corestate_test.py index d43f5afaa86c..2ed03d0c215d 100644 --- a/framework/py/flwr/supercore/corestate/corestate_test.py +++ b/framework/py/flwr/supercore/corestate/corestate_test.py @@ -1720,6 +1720,39 @@ def test_get_task_message_limit(self) -> None: self.assertEqual(pulled[0].metadata.message_id, msg_1.metadata.message_id) self.assertEqual(pulled_next[0].metadata.message_id, msg_2.metadata.message_id) + def test_get_task_message_filters_by_source_task(self) -> None: + """A source-specific claim should not consume another task's message.""" + state = self.state_factory() + run_id = self.task_run_id(state) + src_task_id_1 = state.create_task(task_type=TaskType.MODEL, run_id=run_id) + src_task_id_2 = state.create_task(task_type=TaskType.MODEL, run_id=run_id) + dst_task_id = state.create_task(task_type=TaskType.AGENT_APP, run_id=run_id) + assert ( + src_task_id_1 is not None + and src_task_id_2 is not None + and dst_task_id is not None + ) + reply_1 = create_task_message(src_task_id_1, dst_task_id, run_id) + reply_2 = create_task_message(src_task_id_2, dst_task_id, run_id) + self.assertTrue(state.store_task_message(reply_1)) + self.assertTrue(state.store_task_message(reply_2)) + + pulled_2 = state.get_task_message( + dst_task_ids=[dst_task_id], + src_task_ids=[src_task_id_2], + limit=1, + order_by="created_at", + ) + pulled_1 = state.get_task_message( + dst_task_ids=[dst_task_id], + src_task_ids=[src_task_id_1], + limit=1, + order_by="created_at", + ) + + self.assertEqual(pulled_2[0].metadata.message_id, reply_2.metadata.message_id) + self.assertEqual(pulled_1[0].metadata.message_id, reply_1.metadata.message_id) + def test_store_and_get_task_events(self) -> None: """Task events should round-trip in assigned ID order.""" # Prepare: Create one run with a task and two valid task events. @@ -1749,6 +1782,8 @@ def test_store_and_get_task_events(self) -> None: run_id=run_id, after_task_event_id=events[0].id ) no_new = state.get_task_events(run_id=run_id, after_task_event_id=latest_id) + filtered = state.get_task_events(run_id=run_id, task_ids=[task_id]) + excluded = state.get_task_events(run_id=run_id, task_ids=[task_id + 1]) # Assert: Events keep assigned ID order and cursor filtering works. self.assertEqual(len(events), 2) @@ -1769,6 +1804,8 @@ def test_store_and_get_task_events(self) -> None: self.assertEqual(latest_id, events[1].id) self.assertEqual(after_first, [events[1]]) self.assertEqual(no_new, []) + self.assertEqual(filtered, events) + self.assertEqual(excluded, []) @parameterized.expand( # type: ignore [ diff --git a/framework/py/flwr/supercore/corestate/in_memory_corestate.py b/framework/py/flwr/supercore/corestate/in_memory_corestate.py index f5da8f2f9d94..25012067533c 100644 --- a/framework/py/flwr/supercore/corestate/in_memory_corestate.py +++ b/framework/py/flwr/supercore/corestate/in_memory_corestate.py @@ -1194,6 +1194,7 @@ def get_task_message( self, *, dst_task_ids: Sequence[int] | None = None, + src_task_ids: Sequence[int] | None = None, limit: int | None = None, order_by: Literal["created_at"] | None = None, ) -> Sequence[Message]: @@ -1206,19 +1207,28 @@ def get_task_message( return [] if dst_task_ids is not None and not dst_task_ids: return [] + if src_task_ids is not None and not src_task_ids: + return [] with self.lock_task_store, self.lock_task_message_store: self._cleanup_expired_task_tokens_locked() current = now().timestamp() self._cleanup_invalid_task_messages_locked(current) - # Filter by dst_task_id + # Filter by task IDs dst_task_id_set = set(dst_task_ids) if dst_task_ids is not None else None + src_task_id_set = set(src_task_ids) if src_task_ids is not None else None selected_messages = [ msg for msg in self.task_message_store.values() - if dst_task_id_set is None - or msg.metadata.dst_task_id in dst_task_id_set + if ( + dst_task_id_set is None + or msg.metadata.dst_task_id in dst_task_id_set + ) + and ( + src_task_id_set is None + or msg.metadata.src_task_id in src_task_id_set + ) ] # Apply requested sort order @@ -1264,6 +1274,7 @@ def get_task_events( self, *, run_id: int | None = None, + task_ids: Sequence[int] | None = None, after_task_event_id: int | None = None, ) -> Sequence[TaskEvent]: """Return task-produced run events after the cursor.""" @@ -1277,10 +1288,12 @@ def get_task_events( ] else: events = list(self.task_event_store.get(run_id, [])) + task_id_set = set(task_ids) if task_ids is not None else None return [ event for event in sorted(events, key=lambda event: event.id) if event.id > cursor + and (task_id_set is None or event.task_id in task_id_set) ] def _cleanup_expired_task_tokens_locked(self) -> None: diff --git a/framework/py/flwr/supercore/corestate/sql_corestate.py b/framework/py/flwr/supercore/corestate/sql_corestate.py index 9ee9399404f7..18c3fb19054d 100644 --- a/framework/py/flwr/supercore/corestate/sql_corestate.py +++ b/framework/py/flwr/supercore/corestate/sql_corestate.py @@ -1536,6 +1536,7 @@ def get_task_message( self, *, dst_task_ids: Sequence[int] | None = None, + src_task_ids: Sequence[int] | None = None, limit: int | None = None, order_by: Literal["created_at"] | None = None, ) -> Sequence[Message]: @@ -1548,11 +1549,15 @@ def get_task_message( return [] if dst_task_ids is not None and not dst_task_ids: return [] + if src_task_ids is not None and not src_task_ids: + return [] with self.session(): self._cleanup_expired_task_tokens() self._cleanup_invalid_task_messages() - rows = self._claim_task_message_models(dst_task_ids, order_by, limit) + rows = self._claim_task_message_models( + dst_task_ids, src_task_ids, order_by, limit + ) snapshots = [_task_message_snapshot_from_model(row) for row in rows] return [_task_message_from_snapshot(row) for row in snapshots] @@ -1591,6 +1596,7 @@ def get_task_events( self, *, run_id: int | None = None, + task_ids: Sequence[int] | None = None, after_task_event_id: int | None = None, ) -> Sequence[TaskEvent]: """Return task-produced run events after the cursor.""" @@ -1602,6 +1608,11 @@ def get_task_events( ) if run_id is not None: query = query.where(TaskEventModel.run_id == uint64_to_int64(run_id)) + if task_ids is not None: + if not task_ids: + return [] + sint64_task_ids = [uint64_to_int64(task_id) for task_id in task_ids] + query = query.where(TaskEventModel.task_id.in_(sint64_task_ids)) with self.session() as session: rows = session.scalars(query).all() @@ -1610,6 +1621,7 @@ def get_task_events( def _claim_task_message_models( self, dst_task_ids: Sequence[int] | None, + src_task_ids: Sequence[int] | None, order_by: Literal["created_at"] | None, limit: int | None, ) -> list[TaskMessageModel]: @@ -1620,6 +1632,9 @@ def _claim_task_message_models( if dst_task_ids is not None: sint64_dst_task_ids = [uint64_to_int64(t) for t in dst_task_ids] query = query.where(TaskMessageModel.dst_task_id.in_(sint64_dst_task_ids)) + if src_task_ids is not None: + sint64_src_task_ids = [uint64_to_int64(t) for t in src_task_ids] + query = query.where(TaskMessageModel.src_task_id.in_(sint64_src_task_ids)) if order_by is not None: query = query.order_by(TaskMessageModel.created_at.asc()) if limit is not None: @@ -1646,6 +1661,11 @@ def _claim_task_message_models( delete_query = delete_query.where( TaskMessageModel.dst_task_id.in_(sint64_dst_task_ids) ) + if src_task_ids is not None: + sint64_src_task_ids = [uint64_to_int64(t) for t in src_task_ids] + delete_query = delete_query.where( + TaskMessageModel.src_task_id.in_(sint64_src_task_ids) + ) returning_query = delete_query.returning(TaskMessageModel) with self.session() as session: diff --git a/framework/py/flwr/supercore/json_message/model_message.py b/framework/py/flwr/supercore/json_message/model_message.py index e441149b230a..0f6f5538c479 100644 --- a/framework/py/flwr/supercore/json_message/model_message.py +++ b/framework/py/flwr/supercore/json_message/model_message.py @@ -18,6 +18,7 @@ from __future__ import annotations from collections.abc import Sequence +from typing import cast from flwr.app.constants import DEFAULT_TTL from flwr.supercore.json_message.base import JSONMessage @@ -65,6 +66,24 @@ def __init__( # pylint: disable=too-many-arguments,too-many-positional-argument ttl=ttl, ) + @classmethod + def from_payload(cls, *, dst_task_id: int, payload: JSONObject) -> ModelRequest: + """Create a model request from a Responses request payload.""" + return cls( + dst_task_id=dst_task_id, + input_=cast(str | Sequence[JSONObject], payload.get("input")), + model=cast(str, payload.get("model")), + stream=cast(bool, payload.get("stream", False)), + tools=cast(Sequence[JSONObject] | None, payload.get("tools")), + tool_choice=payload.get("tool_choice"), + reasoning=cast(JSONObject | None, payload.get("reasoning")), + previous_response_id=cast(str | None, payload.get("previous_response_id")), + instructions=cast(str | None, payload.get("instructions")), + max_output_tokens=cast(int | None, payload.get("max_output_tokens")), + metadata=cast(JSONObject | None, payload.get("metadata")), + text=cast(JSONObject | None, payload.get("text")), + ) + @classmethod def _validate_payload(cls, payload: JSONObject) -> None: """Validate the minimal Responses create-request shape.""" diff --git a/framework/py/flwr/supercore/json_message/model_message_test.py b/framework/py/flwr/supercore/json_message/model_message_test.py index db4ba76b6458..537768a49450 100644 --- a/framework/py/flwr/supercore/json_message/model_message_test.py +++ b/framework/py/flwr/supercore/json_message/model_message_test.py @@ -116,10 +116,9 @@ def test_model_messages_create_payloads() -> None: def test_model_request_accepts_string_input_and_default_stream() -> None: """Model requests should accept simple string prompts.""" - request = ModelRequest( + request = ModelRequest.from_payload( dst_task_id=123, - input_="Hello", - model="gpt-5", + payload={"input": "Hello", "model": "gpt-5"}, ) assert request.payload == { diff --git a/framework/py/flwr/supercore/runtime/runtime_http_client.py b/framework/py/flwr/supercore/runtime/runtime_http_client.py index e85bca5abd41..ffc73b2acd9d 100644 --- a/framework/py/flwr/supercore/runtime/runtime_http_client.py +++ b/framework/py/flwr/supercore/runtime/runtime_http_client.py @@ -14,6 +14,13 @@ # ============================================================================== """HTTP client for the Runtime API.""" +from __future__ import annotations + +import json +from typing import NoReturn, cast + +import httpx + from flwr.proto.control_pb2 import ( # pylint: disable=E0611 StartAutomationRequest, StartAutomationResponse, @@ -61,12 +68,47 @@ SendTaskHeartbeatResponse, ) from flwr.supercore.protobuf.client import ProtobufClient +from flwr.supercore.typing import JSONObject + +_TERMINAL_RESPONSE_EVENTS = frozenset( + {"response.completed", "response.failed", "response.incomplete"} +) # Match the method names defined by the Runtime protobuf service. # pylint: disable=invalid-name class RuntimeHttpClient(ProtobufClient): - """Protobuf-over-HTTP client for the Runtime API.""" + """HTTP client for the Runtime API.""" + + def create_response( + self, + request: JSONObject, + *, + token: str, + timeout: float, + ) -> JSONObject: + """Create a model response through the Open Responses endpoint.""" + headers = { + "authorization": f"Bearer {token}", + "accept": ( + "text/event-stream" + if request.get("stream") is True + else "application/json" + ), + } + with self._client.stream( + "POST", + f"{self._base_url}/v1/runtime/responses", + json=request, + headers=headers, + timeout=timeout, + ) as response: + if response.is_error: + response.read() + _raise_for_response_error(response) + if request.get("stream") is True: + return _response_from_stream(response) + return _ensure_json_object(response.json()) def PullPendingTasks( self, request: PullPendingTasksRequest @@ -252,3 +294,56 @@ def GetConnector(self, request: GetConnectorRequest) -> GetConnectorResponse: request=request, response_type=GetConnectorResponse, ) + + +def _response_from_stream(response: httpx.Response) -> JSONObject: + """Return the response object carried by a terminal SSE event.""" + for line in response.iter_lines(): + if not line.startswith("data:"): + continue + data = line.removeprefix("data:").lstrip() + if not data or data == "[DONE]": + continue + try: + event = _ensure_json_object(json.loads(data)) + except json.JSONDecodeError as exc: + raise ValueError("Invalid JSON in Responses event stream.") from exc + + event_type = event.get("type") + if event_type in _TERMINAL_RESPONSE_EVENTS: + return _ensure_json_object(event.get("response")) + if event_type == "error": + error = event.get("error") + return { + "object": "response", + "status": "failed", + "error": error if isinstance(error, dict) else {}, + "output": [], + } + + raise RuntimeError("Responses stream ended without a terminal event.") + + +def _ensure_json_object(payload: object) -> JSONObject: + """Validate and narrow one decoded JSON object.""" + if not isinstance(payload, dict): + raise ValueError("Runtime Responses endpoint returned a non-object payload.") + return cast(JSONObject, payload) + + +def _raise_for_response_error(response: httpx.Response) -> NoReturn: + """Map an Open Responses HTTP error to the AgentResponses API.""" + message = "Runtime Responses request failed." + try: + payload = _ensure_json_object(response.json()) + error = payload.get("error") + if isinstance(error, dict) and isinstance(error.get("message"), str): + message = cast(str, error["message"]) + except (json.JSONDecodeError, ValueError): + pass + + if response.status_code == httpx.codes.BAD_REQUEST: + raise ValueError(message) + if response.status_code == httpx.codes.GATEWAY_TIMEOUT: + raise TimeoutError(message) + raise RuntimeError(message) diff --git a/framework/py/flwr/supercore/runtime/runtime_http_client_test.py b/framework/py/flwr/supercore/runtime/runtime_http_client_test.py index cdeef01f0380..71f4c7af03b8 100644 --- a/framework/py/flwr/supercore/runtime/runtime_http_client_test.py +++ b/framework/py/flwr/supercore/runtime/runtime_http_client_test.py @@ -14,8 +14,10 @@ # ============================================================================== """Tests for the Runtime HTTP client.""" +import json from unittest.mock import Mock, patch +import httpx import pytest from flwr.supercore.protobuf.client import ProtobufClient @@ -71,3 +73,48 @@ def test_runtime_method(endpoint: str) -> None: endpoint, f"{method_name}Response" ) assert call.call_args.kwargs["response_type"].__name__ == expected_response_name + + +@pytest.mark.parametrize("stream", [False, True]) +def test_create_response(stream: bool) -> None: + """Return the final response for JSON and streaming requests.""" + request_payload = {"model": "model", "input": "hello", "stream": stream} + response_payload = { + "object": "response", + "id": "resp-1", + "status": "completed", + "output": [], + } + + def handler(request: httpx.Request) -> httpx.Response: + assert request.url == "http://runtime.example/v1/runtime/responses" + assert request.headers["authorization"] == "Bearer task-token" + assert json.loads(request.content) == request_payload + if stream: + completed_event = { + "type": "response.completed", + "response": response_payload, + } + return httpx.Response( + 200, + headers={"content-type": "text/event-stream"}, + text=( + 'event: response.created\ndata: {"type":"response.created"}\n\n' + "event: response.completed\n" + f"data: {json.dumps(completed_event)}\n\n" + ), + ) + return httpx.Response(200, json=response_payload) + + http_client = httpx.Client(transport=httpx.MockTransport(handler)) + with patch("flwr.supercore.protobuf.client.httpx.Client", return_value=http_client): + client = RuntimeHttpClient("http://runtime.example") + + try: + result = client.create_response( + request_payload, token="task-token", timeout=300.0 + ) + finally: + client.close() + + assert result == response_payload diff --git a/framework/py/flwr/supercore/servicer/runtime/runtime_handlers.py b/framework/py/flwr/supercore/servicer/runtime/runtime_handlers.py index a6844d2b99c8..03f557a04dce 100644 --- a/framework/py/flwr/supercore/servicer/runtime/runtime_handlers.py +++ b/framework/py/flwr/supercore/servicer/runtime/runtime_handlers.py @@ -190,8 +190,10 @@ def pull_task_message( log(DEBUG, "Runtime.PullTaskMessage") limit = request.limit if request.HasField("limit") else None + src_task_ids = [request.src_task_id] if request.HasField("src_task_id") else None messages = state.get_task_message( dst_task_ids=[task.task_id], + src_task_ids=src_task_ids, limit=limit, order_by="created_at", ) diff --git a/framework/py/flwr/supercore/servicer/runtime/runtime_handlers_test.py b/framework/py/flwr/supercore/servicer/runtime/runtime_handlers_test.py index 9300755560f5..af068141de53 100644 --- a/framework/py/flwr/supercore/servicer/runtime/runtime_handlers_test.py +++ b/framework/py/flwr/supercore/servicer/runtime/runtime_handlers_test.py @@ -488,7 +488,7 @@ def test_pull_task_message_uses_authenticated_task_destination(self) -> None: # Execute response = runtime_handlers.pull_task_message( - PullTaskMessageRequest(limit=5), + PullTaskMessageRequest(limit=5, src_task_id=123), self.state, Task(task_id=321, run_id=789), ) @@ -496,6 +496,7 @@ def test_pull_task_message_uses_authenticated_task_destination(self) -> None: # Assert self.state.get_task_message.assert_called_once_with( dst_task_ids=[321], + src_task_ids=[123], limit=5, order_by="created_at", ) diff --git a/framework/py/flwr/supercore/task_process/agent/run_agentapp.py b/framework/py/flwr/supercore/task_process/agent/run_agentapp.py index 47cf9b5f4eca..5a200c6039aa 100644 --- a/framework/py/flwr/supercore/task_process/agent/run_agentapp.py +++ b/framework/py/flwr/supercore/task_process/agent/run_agentapp.py @@ -15,9 +15,12 @@ """Flower AgentApp process.""" +import os +import ssl from logging import DEBUG, ERROR from pathlib import Path from queue import Queue +from tempfile import NamedTemporaryFile import httpx @@ -66,6 +69,10 @@ from .session import RuntimeAgentConnectors, RuntimeAgentResponses, RuntimeAgentSession _AGENT_INPUT_KEY = "agent.input" +_RUNTIME_API_KEY_ENV = "FLWR_RUNTIME_API_KEY" +_RUNTIME_BASE_URL_ENV = "FLWR_RUNTIME_BASE_URL" +_SSL_CERT_FILE_ENV = "SSL_CERT_FILE" +_SSL_CERT_DIR_ENV = "SSL_CERT_DIR" def run_agentapp( # pylint: disable=R0912, R0913, R0914, R0915, R0917, W0212 @@ -97,6 +104,7 @@ def run_agentapp( # pylint: disable=R0912, R0913, R0914, R0915, R0917, W0212 heartbeat_sender = None context: Context | None = None runtime_env_dir: Path | None = None + runtime_root_certificates_path: Path | None = None exit_code = ExitCode.SUCCESS def on_exit() -> None: @@ -126,6 +134,8 @@ def on_exit() -> None: grid.close() cleanup_app_runtime_environment(runtime_env_dir) + if runtime_root_certificates_path is not None: + runtime_root_certificates_path.unlink(missing_ok=True) register_signal_handlers( event_type=EventType.FLWR_AGENTAPP_RUN_LEAVE, @@ -222,6 +232,7 @@ def on_exit() -> None: responses = RuntimeAgentResponses( stub=grid._runtime_client, + token=token, run_id=context.run_id, task_id=task_id, context=context, @@ -236,6 +247,10 @@ def on_exit() -> None: connectors = RuntimeAgentConnectors(responses) agent = RuntimeAgentSession(responses=responses, connectors=connectors) + runtime_root_certificates_path = _set_runtime_environment( + runtime_api_address, token, insecure, certificates + ) + # Load and run the AgentApp agent_app = load_app(agent_app_attr, LoadAgentAppError, app_path) if not isinstance(agent_app, AgentApp): @@ -270,3 +285,67 @@ def on_exit() -> None: "success": exit_code == ExitCode.SUCCESS, }, ) + + +def _set_runtime_environment( + runtime_api_address: str, + token: str, + insecure: bool, + certificates: bytes | None = None, +) -> Path | None: + """Expose the Open Responses-compatible Runtime endpoint to the AgentApp.""" + scheme = "http" if insecure else "https" + address = runtime_api_address.rstrip("/") + os.environ[_RUNTIME_BASE_URL_ENV] = f"{scheme}://{address}/v1/runtime" + os.environ[_RUNTIME_API_KEY_ENV] = token + + if certificates is None: + return None + + # OpenAI/httpx reads custom CAs from SSL_CERT_FILE, which requires a path. + # Extend the public and inherited roots with the Runtime CA for this process. + public_ca_context = httpx.create_ssl_context(trust_env=False) + trusted_certificates = [ + *public_ca_context.get_ca_certs(binary_form=True), + *_load_inherited_ca_certificates(), + ] + with NamedTemporaryFile( + mode="wb", prefix="flwr-runtime-ca-", suffix=".pem", delete=False + ) as certificate_file: + for trusted_certificate in dict.fromkeys(trusted_certificates): + certificate_file.write( + ssl.DER_cert_to_PEM_cert(trusted_certificate).encode("ascii") + ) + certificate_file.write(b"\n") + certificate_file.write(certificates) + certificate_path = Path(certificate_file.name) + os.environ[_SSL_CERT_FILE_ENV] = str(certificate_path) + return certificate_path + + +def _load_inherited_ca_certificates() -> list[bytes]: + """Load roots configured through the standard OpenSSL environment.""" + ca_file = os.environ.get(_SSL_CERT_FILE_ENV) + ca_dir = os.environ.get(_SSL_CERT_DIR_ENV) + if not ca_file and not ca_dir: + return [] + + context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + if ca_file: + context.load_verify_locations(cafile=ca_file) + if ca_dir: + for directory in ca_dir.split(os.pathsep): + if not directory: + continue + try: + ca_paths = sorted(Path(directory).iterdir()) + except OSError: + continue + for ca_path in ca_paths: + if not ca_path.is_file(): + continue + try: + context.load_verify_locations(cafile=ca_path) + except OSError: + continue + return context.get_ca_certs(binary_form=True) diff --git a/framework/py/flwr/supercore/task_process/agent/run_agentapp_test.py b/framework/py/flwr/supercore/task_process/agent/run_agentapp_test.py new file mode 100644 index 000000000000..21e4acef1db1 --- /dev/null +++ b/framework/py/flwr/supercore/task_process/agent/run_agentapp_test.py @@ -0,0 +1,121 @@ +# Copyright 2026 Flower Labs GmbH. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Tests for the AgentApp process environment.""" + +import os +import ssl +from pathlib import Path +from unittest.mock import Mock + +import httpx +import pytest + +from .run_agentapp import _set_runtime_environment + + +@pytest.mark.parametrize(("insecure", "scheme"), [(True, "http"), (False, "https")]) +def test_set_runtime_environment( + monkeypatch: pytest.MonkeyPatch, insecure: bool, scheme: str +) -> None: + """Expose the Runtime Responses base URL and AgentApp task token.""" + monkeypatch.delenv("FLWR_RUNTIME_BASE_URL", raising=False) + monkeypatch.delenv("FLWR_RUNTIME_API_KEY", raising=False) + monkeypatch.delenv("SSL_CERT_FILE", raising=False) + certificate_path = _set_runtime_environment( + "runtime.example:9092", "task-token", insecure=insecure + ) + + assert os.environ["FLWR_RUNTIME_BASE_URL"] == ( + f"{scheme}://runtime.example:9092/v1/runtime" + ) + assert os.environ["FLWR_RUNTIME_API_KEY"] == "task-token" + assert certificate_path is None + assert "SSL_CERT_FILE" not in os.environ + + +def test_set_runtime_environment_exposes_root_certificates( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Add custom Runtime root certificates to the public trust bundle.""" + monkeypatch.delenv("SSL_CERT_FILE", raising=False) + monkeypatch.delenv("SSL_CERT_DIR", raising=False) + + certificate_path = _set_runtime_environment( + "runtime.example:9092", + "task-token", + insecure=False, + certificates=b"root-certificates", + ) + + assert certificate_path is not None + try: + assert os.environ["SSL_CERT_FILE"] == str(certificate_path) + certificate_bundle = certificate_path.read_bytes() + assert certificate_bundle.startswith(b"-----BEGIN CERTIFICATE-----") + assert certificate_bundle.endswith(b"\nroot-certificates") + finally: + certificate_path.unlink(missing_ok=True) + + +@pytest.mark.parametrize("ca_env", ["SSL_CERT_FILE", "SSL_CERT_DIR"]) +def test_set_runtime_environment_preserves_inherited_root_certificates( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, ca_env: str +) -> None: + """Keep inherited custom roots when adding the Runtime root certificate.""" + inherited_certificates = httpx.create_ssl_context(trust_env=False).get_ca_certs( + binary_form=True + )[:2] + inherited_certificates_pem = [ + ssl.DER_cert_to_PEM_cert(certificate).encode("ascii") + for certificate in inherited_certificates + ] + if ca_env == "SSL_CERT_FILE": + inherited_ca_path = tmp_path / "inherited-ca.pem" + inherited_ca_path.write_bytes(inherited_certificates_pem[0]) + ca_env_value = str(inherited_ca_path) + expected_certificates = inherited_certificates_pem[:1] + else: + inherited_ca_directories = [] + for index, certificate in enumerate(inherited_certificates_pem): + inherited_ca_directory = tmp_path / f"inherited-cas-{index}" + inherited_ca_directory.mkdir() + (inherited_ca_directory / "inherited-ca.pem").write_bytes(certificate) + inherited_ca_directories.append(str(inherited_ca_directory)) + ca_env_value = os.pathsep.join(inherited_ca_directories) + expected_certificates = inherited_certificates_pem + monkeypatch.setenv(ca_env, ca_env_value) + monkeypatch.delenv( + "SSL_CERT_DIR" if ca_env == "SSL_CERT_FILE" else "SSL_CERT_FILE", + raising=False, + ) + public_ca_context = Mock() + public_ca_context.get_ca_certs.return_value = [] + monkeypatch.setattr(httpx, "create_ssl_context", lambda **_: public_ca_context) + + certificate_path = _set_runtime_environment( + "runtime.example:9092", + "task-token", + insecure=False, + certificates=b"runtime-root-certificate", + ) + + assert certificate_path is not None + try: + certificate_bundle = certificate_path.read_bytes() + for certificate in expected_certificates: + assert certificate in certificate_bundle + assert certificate_bundle.endswith(b"\nruntime-root-certificate") + finally: + certificate_path.unlink(missing_ok=True) diff --git a/framework/py/flwr/supercore/task_process/agent/session.py b/framework/py/flwr/supercore/task_process/agent/session.py index 62667d98f50b..fff5c3a0c2b8 100644 --- a/framework/py/flwr/supercore/task_process/agent/session.py +++ b/framework/py/flwr/supercore/task_process/agent/session.py @@ -43,7 +43,6 @@ ConnectorRequest, ConnectorResponse, ) -from flwr.supercore.json_message.model_message import ModelRequest, ModelResponse from flwr.supercore.runtime import RuntimeHttpClient from flwr.supercore.task_process.connector.automation import START_AUTOMATION_TOOL_NAME from flwr.supercore.task_process.connector.registry import ( @@ -112,18 +111,20 @@ def call(self, tool_call: JSONObject) -> JSONObject: class RuntimeAgentResponses(AgentResponses): - """AgentResponses implementation backed by Runtime task messages.""" + """AgentResponses implementation backed by the Runtime API.""" def __init__( # pylint: disable=too-many-arguments self, *, stub: RuntimeHttpClient, + token: str, run_id: int, task_id: int, context: Context, start_run_request: StartRunRequest, ) -> None: self._stub = stub + self._token = token self._context = context self._run_id = run_id self._task_id = task_id @@ -140,36 +141,11 @@ def create(self, request: JSONObject) -> JSONObject: def _create_model_response(self, request: JSONObject) -> JSONObject: """Create one model response through a child model task.""" - model = request.get("model") - if not isinstance(model, str) or not model: - raise ValueError( - "AgentResponses request requires a non-empty string 'model' field." - ) - - create_res = self._stub.CreateTask( - CreateTaskRequest(type=TaskType.MODEL, model_ref=model) - ) - if not create_res.HasField("task_id"): - raise RuntimeError("Model task could not be created.") - - model_task_id = create_res.task_id - message = ModelRequest( - dst_task_id=model_task_id, - input_=cast(str | Sequence[JSONObject], request.get("input")), - model=model, - stream=cast(bool, request.get("stream", False)), - tools=cast(Sequence[JSONObject] | None, request.get("tools")), - tool_choice=request.get("tool_choice"), - reasoning=cast(JSONObject | None, request.get("reasoning")), - previous_response_id=cast(str | None, request.get("previous_response_id")), - instructions=cast(str | None, request.get("instructions")), - max_output_tokens=cast(int | None, request.get("max_output_tokens")), - metadata=cast(JSONObject | None, request.get("metadata")), - text=cast(JSONObject | None, request.get("text")), + return self._stub.create_response( + request, + token=self._token, + timeout=_DEFAULT_MODEL_REPLY_TIMEOUT, ) - response_message = self._send_and_receive(message) - response = ModelResponse.from_message(response_message) - return response.payload def create_connector_response( self, *, name: str, call_id: str, arguments: JSONObject @@ -359,9 +335,11 @@ def _push_task_message(self, message: Message) -> None: PushTaskMessageRequest(message=message_to_proto(message)) ) - def _pull_task_messages(self) -> list[Message]: - """Pull pending task messages.""" - res = self._stub.PullTaskMessage(PullTaskMessageRequest(limit=1)) + def _pull_task_messages(self, src_task_id: int) -> list[Message]: + """Pull pending task messages from one child task.""" + res = self._stub.PullTaskMessage( + PullTaskMessageRequest(limit=1, src_task_id=src_task_id) + ) return [message_from_proto(msg) for msg in res.messages] def _send_and_receive(self, message: Message) -> Message: @@ -370,6 +348,10 @@ def _send_and_receive(self, message: Message) -> Message: For now, `flwr-agentapp` expects a strict one-request-one-reply exchange with child tasks, so any non-matching pulled message is treated as an error. """ + child_task_id = message.metadata.dst_task_id + if child_task_id is None: + raise ValueError("Task message requires a destination task ID.") + # Push the message to the child task self._push_task_message(message) message_id = message.metadata.message_id @@ -377,7 +359,7 @@ def _send_and_receive(self, message: Message) -> Message: # Pull until a message arrives that replies to the pushed message, or timeout deadline = time.monotonic() + _DEFAULT_MODEL_REPLY_TIMEOUT while True: - for pulled_msg in self._pull_task_messages(): + for pulled_msg in self._pull_task_messages(child_task_id): if pulled_msg.metadata.reply_to_message_id != message_id: raise RuntimeError( "Received a message that does not reply to the request." diff --git a/framework/py/flwr/supercore/task_process/agent/session_test.py b/framework/py/flwr/supercore/task_process/agent/session_test.py index 42aeea53c34e..b983e3c1b87d 100644 --- a/framework/py/flwr/supercore/task_process/agent/session_test.py +++ b/framework/py/flwr/supercore/task_process/agent/session_test.py @@ -26,6 +26,8 @@ from flwr.proto.runtime_pb2 import ( # pylint: disable=E0611 CreateTaskRequest, CreateTaskResponse, + PullTaskMessageRequest, + PullTaskMessageResponse, ) from flwr.supercore.constant import TaskType from flwr.supercore.json_message.connector_message import ( @@ -39,6 +41,25 @@ from .session import RuntimeAgentConnectors, RuntimeAgentResponses +def test_pull_task_messages_filters_by_child_task() -> None: + """Claim only messages sent by the expected child task.""" + stub = Mock() + stub.PullTaskMessage.return_value = PullTaskMessageResponse() + responses = RuntimeAgentResponses( + stub=stub, + token="task-token", + run_id=123, + task_id=789, + context=Mock(), + start_run_request=StartRunRequest(), + ) + + assert responses._pull_task_messages(456) == [] # pylint: disable=W0212 + stub.PullTaskMessage.assert_called_once_with( + PullTaskMessageRequest(limit=1, src_task_id=456) + ) + + def test_start_automation_tool_exposes_only_input_and_schedule() -> None: """Keep the embedded run request out of the model-facing schema.""" # Prepare @@ -72,6 +93,38 @@ def test_runtime_connectors_expand_one_connector_into_multiple_tools() -> None: get_connector_tools.assert_called_once_with("example") +def test_create_response_delegates_to_runtime_endpoint() -> None: + """Delegate model exchange handling and retain response context output.""" + stub = Mock() + response: JSONObject = { + "object": "response", + "status": "completed", + "output": [{"type": "message", "role": "assistant"}], + } + stub.create_response.return_value = response + context = Mock() + responses = RuntimeAgentResponses( + stub=stub, + token="task-token", + run_id=123, + task_id=789, + context=context, + start_run_request=StartRunRequest(), + ) + request: JSONObject = {"model": "model", "input": "hello", "stream": True} + + with patch( + "flwr.supercore.task_process.agent.session.append_items" + ) as append_items: + result = responses.create(request) + + assert result == response + stub.create_response.assert_called_once_with( + request, token="task-token", timeout=300.0 + ) + append_items.assert_called_once_with(context, response["output"]) + + def test_call_automation_embeds_input_in_control_request() -> None: """Embed model input in the Control request sent to the Runtime API.""" # Prepare @@ -85,6 +138,7 @@ def test_call_automation_embeds_input_in_control_request() -> None: ) responses = RuntimeAgentResponses( stub=stub, + token="task-token", run_id=123, task_id=789, context=Mock(), @@ -127,6 +181,7 @@ def test_create_connector_response_resolves_canonical_name() -> None: stub.CreateTask.return_value = CreateTaskResponse(task_id=456) responses = RuntimeAgentResponses( stub=stub, + token="task-token", run_id=123, task_id=789, context=Mock(), diff --git a/framework/py/flwr/superlink/main.py b/framework/py/flwr/superlink/main.py index 10a83ddfa623..f8df4a0972d1 100644 --- a/framework/py/flwr/superlink/main.py +++ b/framework/py/flwr/superlink/main.py @@ -50,6 +50,7 @@ ControlEventLogMiddleware, ControlLicenseMiddleware, ) +from flwr.superlink.routers.runtime import responses_router from flwr.superlink.routers.runtime import router as runtime_router try: @@ -184,6 +185,7 @@ async def lifespan(fastapi_app: FastAPI) -> AsyncIterator[dict[str, object]]: # SuperLink APIs fastapi_app.include_router(control_router) fastapi_app.include_router(runtime_router) + fastapi_app.include_router(responses_router) # Extension hooks extensions.configure_app(fastapi_app) diff --git a/framework/py/flwr/superlink/routers/runtime/__init__.py b/framework/py/flwr/superlink/routers/runtime/__init__.py index 69272131b49e..7dcb7ede779b 100644 --- a/framework/py/flwr/superlink/routers/runtime/__init__.py +++ b/framework/py/flwr/superlink/routers/runtime/__init__.py @@ -15,6 +15,7 @@ """Runtime API router.""" +from .responses import router as responses_router from .router import router -__all__ = ["router"] +__all__ = ["responses_router", "router"] diff --git a/framework/py/flwr/superlink/routers/runtime/responses.py b/framework/py/flwr/superlink/routers/runtime/responses.py new file mode 100644 index 000000000000..1b31d1707c92 --- /dev/null +++ b/framework/py/flwr/superlink/routers/runtime/responses.py @@ -0,0 +1,437 @@ +# Copyright 2026 Flower Labs GmbH. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Open Responses-compatible Runtime endpoint.""" + +from __future__ import annotations + +import asyncio +import json +import os +import time +from collections.abc import AsyncGenerator +from dataclasses import dataclass +from typing import Annotated, cast + +from fastapi import APIRouter, Depends, Request, Response +from fastapi.responses import JSONResponse, StreamingResponse +from starlette.concurrency import run_in_threadpool + +from flwr.common.constant import Status, SubStatus +from flwr.proto.runtime_pb2 import CreateTaskRequest # pylint: disable=E0611 +from flwr.proto.task_pb2 import Task, TaskEvent # pylint: disable=E0611 +from flwr.server.superlink.linkstate import LinkState +from flwr.supercore.constant import TaskType +from flwr.supercore.error import FlowerError +from flwr.supercore.json_message.model_message import ModelRequest, ModelResponse +from flwr.supercore.servicer.runtime import runtime_handlers +from flwr.supercore.typing import JSONObject +from flwr.supercore.utils import strict_json_dumps +from flwr.superlink.dependencies.linkstate import get_linkstate + +router = APIRouter(prefix="/v1/runtime", tags=["Runtime"]) + +LinkStateDependency = Annotated[LinkState, Depends(get_linkstate)] + +_SUPPORTED_FIELDS = frozenset( + { + "model", + "input", + "stream", + "tools", + "tool_choice", + "reasoning", + "previous_response_id", + "instructions", + "max_output_tokens", + "metadata", + "text", + } +) +_TERMINAL_EVENTS = frozenset( + {"error", "response.completed", "response.failed", "response.incomplete"} +) +_POLL_INTERVAL = 0.25 +_DEFAULT_MODEL_TASK_LAUNCH_TIMEOUT = 300.0 +_MODEL_TASK_LAUNCH_TIMEOUT_ENV = "FLWR_MODEL_TASK_LAUNCH_TIMEOUT" + + +@dataclass(frozen=True) +class _Exchange: + """Identify one AgentApp-to-model task exchange.""" + + agent_task_id: int + model_task_id: int + run_id: int + + +class _ResponsesError(Exception): + """Represent an error returned through the Responses HTTP contract.""" + + def __init__(self, status_code: int, message: str, code: str) -> None: + super().__init__(message) + self.status_code = status_code + self.message = message + self.code = code + + +@router.post("/responses") +async def create_runtime_response( + request: Request, + state: LinkStateDependency, +) -> Response: + """Create a model response through a child model task.""" + try: + task = await run_in_threadpool(_authenticate, request, state) + payload = await _read_request_payload(request) + model_request = _model_request_from_payload(payload) + exchange = await run_in_threadpool(_start_exchange, state, task, model_request) + except _ResponsesError as err: + return _error_response(err) + + if model_request.payload.get("stream") is True: + return StreamingResponse( + _stream_response(state, exchange), + media_type="text/event-stream", + headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}, + ) + + try: + response = await _wait_for_response(request, state, exchange) + except _ResponsesError as err: + return _error_response(err) + return JSONResponse(content=response) + + +def _authenticate(request: Request, state: LinkState) -> Task: + """Authenticate exactly one AgentApp Bearer token.""" + authorization = request.headers.getlist("authorization") + if len(authorization) != 1: + raise _ResponsesError( + 401, "Invalid authentication credentials.", "invalid_api_key" + ) + + parts = authorization[0].split() + if len(parts) != 2 or parts[0].lower() != "bearer" or not parts[1]: + raise _ResponsesError( + 401, "Invalid authentication credentials.", "invalid_api_key" + ) + + task = state.get_task_by_token(parts[1]) + if task is None or task.type != TaskType.AGENT_APP: + raise _ResponsesError( + 401, "Invalid authentication credentials.", "invalid_api_key" + ) + return task + + +async def _read_request_payload(request: Request) -> JSONObject: + """Parse and validate the JSON request object.""" + try: + payload = await request.json() + except (json.JSONDecodeError, UnicodeDecodeError) as err: + raise _ResponsesError( + 400, "Request body must be valid JSON.", "invalid_json" + ) from err + + if not isinstance(payload, dict): + raise _ResponsesError( + 400, "Request body must be a JSON object.", "invalid_request" + ) + unsupported = sorted(set(payload) - _SUPPORTED_FIELDS) + if unsupported: + fields = ", ".join(unsupported) + raise _ResponsesError( + 400, + f"Unsupported request field(s): {fields}.", + "unsupported_parameter", + ) + return cast(JSONObject, payload) + + +def _model_request_from_payload(payload: JSONObject) -> ModelRequest: + """Build and validate the task-routed model request.""" + try: + return ModelRequest.from_payload(dst_task_id=0, payload=payload) + except (TypeError, ValueError) as err: + raise _ResponsesError(400, str(err), "invalid_request") from err + + +def _start_exchange( + state: LinkState, + task: Task, + request: ModelRequest, +) -> _Exchange: + """Create a child model task and send its request message.""" + model = cast(str, request.payload["model"]) + try: + response = runtime_handlers.create_task( + CreateTaskRequest(type=TaskType.MODEL, model_ref=model), state, task + ) + except FlowerError as err: + raise _ResponsesError( + 500, "Model task could not be created.", "model_task_creation_failed" + ) from err + if not response.HasField("task_id"): + raise _ResponsesError( + 500, "Model task could not be created.", "model_task_creation_failed" + ) + + model_task_id = response.task_id + request.metadata.dst_task_id = model_task_id + request.metadata.__dict__["_run_id"] = task.run_id + request.metadata.src_task_id = task.task_id + request.metadata.__dict__["_message_id"] = request.object_id + if not state.store_task_message(request): + state.finish_task( + model_task_id, SubStatus.STOPPED, "Model request was not stored." + ) + raise _ResponsesError( + 500, "Model request could not be stored.", "model_request_failed" + ) + + return _Exchange( + agent_task_id=task.task_id, + model_task_id=model_task_id, + run_id=task.run_id, + ) + + +async def _wait_for_response( + request: Request, state: LinkState, exchange: _Exchange +) -> JSONObject: + """Wait for and return one correlated model response.""" + launch_deadline = time.monotonic() + _model_task_launch_timeout() + complete = False + try: + while True: + if await request.is_disconnected(): + raise _ResponsesError( + 499, "Client disconnected.", "client_disconnected" + ) + + response = await run_in_threadpool( + _claim_response_or_raise_for_model_task_state, + state, + exchange, + launch_deadline, + ) + if response is not None: + complete = True + _raise_for_failed_response(response) + return response + + await asyncio.sleep(_POLL_INTERVAL) + finally: + if not complete: + await run_in_threadpool( + _stop_model_task, state, exchange, "Responses request ended early." + ) + + +async def _stream_response( + state: LinkState, exchange: _Exchange +) -> AsyncGenerator[str, None]: + """Relay correlated child-task events as Server-Sent Events.""" + launch_deadline = time.monotonic() + _model_task_launch_timeout() + cursor: int | None = None + complete = False + response: JSONObject | None = None + try: + while True: + events = await run_in_threadpool( + state.get_task_events, + run_id=exchange.run_id, + task_ids=[exchange.model_task_id], + after_task_event_id=cursor, + ) + for event in events: + cursor = event.id + if event.event in _TERMINAL_EVENTS: + if response is None: + await _wait_for_terminal_reply(state, exchange) + complete = True + yield _sse_frame(event) + return + yield _sse_frame(event) + + if response is not None: + complete = True + yield _stream_error( + ( + _response_error_message(response) + if response.get("status") == "failed" + else "Model stream ended before a terminal event." + ), + "model_stream_failed", + ) + return + + response = await run_in_threadpool( + _claim_response_or_raise_for_model_task_state, + state, + exchange, + launch_deadline, + ) + if response is None: + await asyncio.sleep(_POLL_INTERVAL) + except _ResponsesError as err: + yield _stream_error(err.message, err.code) + finally: + if not complete: + await run_in_threadpool( + _stop_model_task, state, exchange, "Responses stream ended early." + ) + + +async def _wait_for_terminal_reply( + state: LinkState, + exchange: _Exchange, +) -> JSONObject: + """Consume the final reply before exposing a terminal stream event.""" + while True: + response = await run_in_threadpool( + _claim_response_or_raise_for_model_task_state, state, exchange + ) + if response is not None: + return response + await asyncio.sleep(_POLL_INTERVAL) + + +def _claim_response(state: LinkState, exchange: _Exchange) -> JSONObject | None: + """Atomically claim the reply belonging to this exchange.""" + messages = state.get_task_message( + dst_task_ids=[exchange.agent_task_id], + src_task_ids=[exchange.model_task_id], + limit=1, + order_by="created_at", + ) + if not messages: + return None + try: + return ModelResponse.from_message(messages[0]).payload + except ValueError as err: + raise _ResponsesError( + 502, "Model task returned an invalid response.", "invalid_model_response" + ) from err + + +def _claim_response_or_raise_for_model_task_state( + state: LinkState, + exchange: _Exchange, + launch_deadline: float | None = None, +) -> JSONObject | None: + """Claim a reply, or fail when its task cannot produce one.""" + response = _claim_response(state, exchange) + if response is not None: + return response + + tasks = state.get_tasks(task_ids=[exchange.model_task_id]) + if tasks and tasks[0].status.status == Status.FINISHED: + # The reply is stored immediately before the task is marked finished. + response = _claim_response(state, exchange) + if response is not None: + return response + if not tasks or tasks[0].status.status == Status.FINISHED: + details = tasks[0].status.details if tasks else "" + raise _ResponsesError( + 502, + details or "Model task ended without a response.", + "model_task_failed", + ) + if ( + launch_deadline is not None + and tasks[0].status.status in {Status.PENDING, Status.STARTING} + and time.monotonic() >= launch_deadline + ): + raise _ResponsesError( + 504, + "Model task was not launched before the configured timeout.", + "model_task_launch_timeout", + ) + return None + + +def _model_task_launch_timeout() -> float: + """Return the configured maximum time for launching a model task.""" + raw_timeout = os.getenv( + _MODEL_TASK_LAUNCH_TIMEOUT_ENV, + str(_DEFAULT_MODEL_TASK_LAUNCH_TIMEOUT), + ) + try: + timeout = float(raw_timeout.strip()) + except ValueError: + timeout = _DEFAULT_MODEL_TASK_LAUNCH_TIMEOUT + return max(1.0, timeout) + + +def _raise_for_failed_response(response: JSONObject) -> None: + """Map a structured failed ModelResponse to an HTTP error.""" + if response.get("status") == "failed" or response.get("error") is not None: + raise _ResponsesError( + 502, _response_error_message(response), "model_provider_error" + ) + + +def _response_error_message(response: JSONObject) -> str: + """Extract a safe message from a structured failed response.""" + error = response.get("error") + if isinstance(error, dict) and isinstance(error.get("message"), str): + return cast(str, error["message"]) + return "Model request failed." + + +def _stop_model_task(state: LinkState, exchange: _Exchange, details: str) -> None: + """Stop an unfinished model task and drain an already-arrived reply.""" + state.finish_task(exchange.model_task_id, SubStatus.STOPPED, details) + try: + _claim_response(state, exchange) + except _ResponsesError: + # Cleanup is best-effort; the original request has already ended. + pass + + +def _sse_frame(event: TaskEvent) -> str: + """Encode one stored task event as an SSE frame.""" + return f"event: {event.event}\ndata: {event.data}\n\n" + + +def _stream_error(message: str, code: str) -> str: + """Encode a terminal Responses error event.""" + data: JSONObject = { + "type": "error", + "error": {"type": "server_error", "code": code, "message": message}, + } + return f"event: error\ndata: {strict_json_dumps(data, compact=True)}\n\n" + + +def _error_response(error: _ResponsesError) -> JSONResponse: + """Return an OpenAI-compatible JSON error envelope.""" + headers = {"WWW-Authenticate": "Bearer"} if error.status_code == 401 else None + return JSONResponse( + status_code=error.status_code, + headers=headers, + content={ + "error": { + "message": error.message, + "type": ( + "invalid_request_error" + if error.status_code in {400, 401} + else "server_error" + ), + "param": None, + "code": error.code, + } + }, + ) diff --git a/framework/py/flwr/superlink/routers/runtime/responses_test.py b/framework/py/flwr/superlink/routers/runtime/responses_test.py new file mode 100644 index 000000000000..9cdc82317128 --- /dev/null +++ b/framework/py/flwr/superlink/routers/runtime/responses_test.py @@ -0,0 +1,358 @@ +# Copyright 2026 Flower Labs GmbH. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Tests for the Runtime Responses endpoint.""" + +import asyncio +from unittest.mock import AsyncMock, Mock, patch + +import pytest +from fastapi import FastAPI, Request +from fastapi.testclient import TestClient + +from flwr.common.constant import Status, SubStatus +from flwr.proto.task_pb2 import Task, TaskEvent, TaskStatus # pylint: disable=E0611 +from flwr.server.superlink.linkstate import LinkState +from flwr.supercore.constant import TaskType +from flwr.supercore.json_message.model_message import ModelResponse +from flwr.superlink.dependencies.linkstate import get_linkstate + +from .responses import ( + _Exchange, + _ResponsesError, + _stream_response, + _wait_for_response, + router, +) + + +def _client(state: Mock) -> TestClient: + app = FastAPI() + app.include_router(router) + app.dependency_overrides[get_linkstate] = lambda: state + return TestClient(app) + + +def _state() -> Mock: + state = Mock(spec=LinkState) + state.get_task_by_token.return_value = Task( + task_id=123, run_id=789, type=TaskType.AGENT_APP + ) + state.create_task.return_value = 456 + state.store_task_message.return_value = True + return state + + +def _reply(request_message_id: str) -> ModelResponse: + return ModelResponse( + dst_task_id=123, + response={ + "object": "response", + "id": "resp_1", + "status": "completed", + "output": [], + }, + reply_to_message_id=request_message_id, + ) + + +def _event(event_id: int, event: str, data: str | None = None) -> TaskEvent: + """Create one child model task event.""" + return TaskEvent( + id=event_id, + run_id=789, + task_id=456, + event=event, + data=data or f'{{"type":"{event}"}}', + ) + + +@pytest.mark.parametrize("authorization", [None, "Basic task-token"]) +def test_responses_requires_bearer_authentication( + authorization: str | None, +) -> None: + """Reject missing and non-Bearer task credentials.""" + headers = {"Authorization": authorization} if authorization else {} + + response = _client(_state()).post( + "/v1/runtime/responses", + json={"model": "model", "input": "hello"}, + headers=headers, + ) + + assert response.status_code == 401 + assert response.json()["error"]["code"] == "invalid_api_key" + + +def test_responses_returns_correlated_model_response() -> None: + """Recheck and claim the direct reply after its child task finishes.""" + state = _state() + state.get_task_message.side_effect = [ + [], + [_reply("request-message-id")], + ] + state.get_tasks.return_value = [ + Task(task_id=456, status=TaskStatus(status=Status.FINISHED)) + ] + + response = _client(state).post( + "/v1/runtime/responses", + json={"model": "model", "input": "hello"}, + headers={"Authorization": "Bearer task-token"}, + ) + + assert response.status_code == 200 + assert response.json()["id"] == "resp_1" + request = state.store_task_message.call_args.args[0] + assert request.metadata.src_task_id == 123 + assert request.metadata.dst_task_id == 456 + assert state.get_task_message.call_count == 2 + state.get_task_message.assert_called_with( + dst_task_ids=[123], + src_task_ids=[456], + limit=1, + order_by="created_at", + ) + + +def test_responses_rejects_unsupported_fields() -> None: + """Do not silently discard unsupported Open Responses fields.""" + state = _state() + + response = _client(state).post( + "/v1/runtime/responses", + json={"model": "model", "input": "hello", "temperature": 0.5}, + headers={"Authorization": "Bearer task-token"}, + ) + + assert response.status_code == 400 + assert response.json()["error"]["code"] == "unsupported_parameter" + state.create_task.assert_not_called() + + +def test_responses_streams_only_child_task_events() -> None: + """Relay ordered events and consume the final correlated reply.""" + state = _state() + state.get_task_events.return_value = [ + _event(1, "response.created"), + _event(2, "response.completed"), + ] + state.get_task_message.side_effect = lambda **_: [ + _reply(state.store_task_message.call_args.args[0].metadata.message_id) + ] + + response = _client(state).post( + "/v1/runtime/responses", + json={"model": "model", "input": "hello", "stream": True}, + headers={"Authorization": "Bearer task-token"}, + ) + + assert response.status_code == 200 + assert response.headers["content-type"].startswith("text/event-stream") + assert response.text == ( + 'event: response.created\ndata: {"type":"response.created"}\n\n' + 'event: response.completed\ndata: {"type":"response.completed"}\n\n' + ) + state.get_task_events.assert_called_once_with( + run_id=789, task_ids=[456], after_task_event_id=None + ) + + +def test_responses_waits_for_terminal_events_after_reply() -> None: + """Poll once more for the terminal event after observing the final reply.""" + state = _state() + state.get_task_events.side_effect = [ + [ + _event( + 1, + "response.output_text.delta", + '{"type":"response.output_text.delta","delta":"x"}', + ) + ], + [_event(2, "response.completed")], + ] + state.get_task_message.side_effect = lambda **_: [ + _reply(state.store_task_message.call_args.args[0].metadata.message_id) + ] + + response = _client(state).post( + "/v1/runtime/responses", + json={"model": "model", "input": "hello", "stream": True}, + headers={"Authorization": "Bearer task-token"}, + ) + + assert "event: response.output_text.delta" in response.text + assert "event: response.completed" in response.text + assert "event: error" not in response.text + assert state.get_task_events.call_count == 2 + + +def test_responses_reports_missing_terminal_event_after_reply() -> None: + """Fail when the flushed event stream has no terminal event.""" + state = _state() + state.get_task_events.return_value = [] + state.get_task_message.side_effect = lambda **_: [ + _reply(state.store_task_message.call_args.args[0].metadata.message_id) + ] + + response = _client(state).post( + "/v1/runtime/responses", + json={"model": "model", "input": "hello", "stream": True}, + headers={"Authorization": "Bearer task-token"}, + ) + + assert "event: error" in response.text + assert "Model stream ended before a terminal event." in response.text + assert state.get_task_events.call_count == 2 + + +def test_responses_stops_and_drains_when_response_wait_is_cancelled() -> None: + """Stop the child task and drain its reply when response waiting is cancelled.""" + state = _state() + state.get_task_message.return_value = [] + state.get_tasks.return_value = [Task(task_id=456)] + + async def wait_until_polled() -> None: + while not state.get_tasks.called: + await asyncio.sleep(0) + + async def cancel_response_wait() -> None: + request = Mock(spec=Request) + request.is_disconnected = AsyncMock(return_value=False) + exchange = _Exchange( + agent_task_id=123, + model_task_id=456, + run_id=789, + ) + with patch("flwr.superlink.routers.runtime.responses._POLL_INTERVAL", new=10): + response_wait = asyncio.create_task( + _wait_for_response(request, state, exchange) + ) + await asyncio.wait_for(wait_until_polled(), timeout=1) + response_wait.cancel() + with pytest.raises(asyncio.CancelledError): + await response_wait + + asyncio.run(cancel_response_wait()) + + state.finish_task.assert_called_once_with( + 456, SubStatus.STOPPED, "Responses request ended early." + ) + assert state.get_task_message.call_count == 2 + + +def test_responses_stops_and_drains_when_client_disconnects() -> None: + """Stop the child task and drain its reply when the client disconnects.""" + state = _state() + state.get_task_message.return_value = [] + request = Mock(spec=Request) + request.is_disconnected = AsyncMock(return_value=True) + + async def wait_for_response() -> None: + exchange = _Exchange( + agent_task_id=123, + model_task_id=456, + run_id=789, + ) + with pytest.raises(_ResponsesError) as exc_info: + await _wait_for_response(request, state, exchange) + assert exc_info.value.code == "client_disconnected" + + asyncio.run(wait_for_response()) + + state.finish_task.assert_called_once_with( + 456, SubStatus.STOPPED, "Responses request ended early." + ) + state.get_task_message.assert_called_once_with( + dst_task_ids=[123], + src_task_ids=[456], + limit=1, + order_by="created_at", + ) + + +def test_responses_times_out_if_model_task_is_not_launched() -> None: + """Stop waiting when no executor launches the child model task.""" + state = _state() + state.get_task_message.return_value = [] + state.get_tasks.return_value = [ + Task( + task_id=456, + status=TaskStatus(status=Status.PENDING), + ) + ] + request = Mock(spec=Request) + request.is_disconnected = AsyncMock(return_value=False) + + async def wait_for_response() -> None: + exchange = _Exchange( + agent_task_id=123, + model_task_id=456, + run_id=789, + ) + with patch( + "flwr.superlink.routers.runtime.responses._model_task_launch_timeout", + return_value=0.0, + ): + with pytest.raises(_ResponsesError) as exc_info: + await _wait_for_response(request, state, exchange) + assert exc_info.value.status_code == 504 + assert exc_info.value.code == "model_task_launch_timeout" + + asyncio.run(wait_for_response()) + + state.finish_task.assert_called_once_with( + 456, SubStatus.STOPPED, "Responses request ended early." + ) + + +def test_responses_stops_and_drains_when_stream_is_cancelled() -> None: + """Stop the child task and drain its reply when the stream is cancelled.""" + state = _state() + state.get_task_events.return_value = [] + state.get_task_message.return_value = [] + state.get_tasks.return_value = [Task(task_id=456)] + + async def wait_until_polled() -> None: + while not state.get_tasks.called: + await asyncio.sleep(0) + + async def cancel_stream() -> None: + stream = _stream_response( + state, + _Exchange( + agent_task_id=123, + model_task_id=456, + run_id=789, + ), + ) + with patch("flwr.superlink.routers.runtime.responses._POLL_INTERVAL", new=10): + next_event = asyncio.create_task(anext(stream)) + await asyncio.wait_for(wait_until_polled(), timeout=1) + next_event.cancel() + with pytest.raises(asyncio.CancelledError): + await next_event + + asyncio.run(cancel_stream()) + + state.finish_task.assert_called_once_with( + 456, SubStatus.STOPPED, "Responses stream ended early." + ) + assert state.get_task_message.call_count == 2 + state.get_task_message.assert_called_with( + dst_task_ids=[123], + src_task_ids=[456], + limit=1, + order_by="created_at", + )