fix: harden streaming transport and ownership checks

This commit is contained in:
a.tolmachev
2026-04-11 01:29:01 +03:00
parent 1a4f0ea6f3
commit 2770c5935f
3 changed files with 348 additions and 7 deletions
+107 -3
View File
@@ -889,6 +889,16 @@ async fn handle_session_poll_call(
Err(error) => return internal_jsonrpc_error(message, error),
};
if !stream_session_belongs_to_tool(&loaded, &tool) {
return tool_error_response(
message,
response_mode,
&session.protocol_version,
"stream_session_not_found",
format!("stream session {} was not found", control.session_id),
);
}
if loaded.is_expired(&now) {
let _ = state
.registry
@@ -1008,7 +1018,7 @@ async fn handle_session_stop_call(
session: &SessionState,
message: &Value,
response_mode: ResponseMode,
_tool: PublishedAgentTool,
tool: PublishedAgentTool,
arguments: Value,
) -> Response {
let control: SessionControlArgs = match serde_json::from_value(arguments.clone()) {
@@ -1024,6 +1034,34 @@ async fn handle_session_stop_call(
}
};
let loaded = match state
.registry
.get_stream_session(&StreamSessionId::new(control.session_id.clone()))
.await
{
Ok(Some(session_record)) => session_record,
Ok(None) => {
return tool_error_response(
message,
response_mode,
&session.protocol_version,
"stream_session_not_found",
format!("stream session {} was not found", control.session_id),
);
}
Err(error) => return internal_jsonrpc_error(message, error),
};
if !stream_session_belongs_to_tool(&loaded, &tool) {
return tool_error_response(
message,
response_mode,
&session.protocol_version,
"stream_session_not_found",
format!("stream session {} was not found", control.session_id),
);
}
match state
.registry
.close_stream_session(
@@ -1075,7 +1113,10 @@ async fn handle_async_job_start_call(
progress: json!({ "pct": 0 }),
result: None,
error: None,
expires_at: Some(add_millis(&now, 300_000)),
expires_at: Some(add_millis(
&now,
streaming.max_session_lifetime_ms.unwrap_or(300_000),
)),
created_at: now.clone(),
updated_at: now.clone(),
finished_at: None,
@@ -1222,6 +1263,16 @@ async fn handle_async_job_status_call(
Err(error) => return internal_jsonrpc_error(message, error),
};
if !async_job_belongs_to_tool(&job, &_tool) {
return tool_error_response(
message,
response_mode,
&session.protocol_version,
"async_job_not_found",
format!("async job {} was not found", control.job_id),
);
}
success_tool_response(
message,
response_mode,
@@ -1275,6 +1326,16 @@ async fn handle_async_job_result_call(
Err(error) => return internal_jsonrpc_error(message, error),
};
if !async_job_belongs_to_tool(&job, &_tool) {
return tool_error_response(
message,
response_mode,
&session.protocol_version,
"async_job_not_found",
format!("async job {} was not found", control.job_id),
);
}
match job.status {
JobStatus::Completed => success_tool_response(
message,
@@ -1330,6 +1391,34 @@ async fn handle_async_job_cancel_call(
}
};
match state
.registry
.get_async_job(&AsyncJobId::new(control.job_id.clone()))
.await
{
Ok(Some(job)) => {
if !async_job_belongs_to_tool(&job, &_tool) {
return tool_error_response(
message,
response_mode,
&session.protocol_version,
"async_job_not_found",
format!("async job {} was not found", control.job_id),
);
}
}
Ok(None) => {
return tool_error_response(
message,
response_mode,
&session.protocol_version,
"async_job_not_found",
format!("async job {} was not found", control.job_id),
);
}
Err(error) => return internal_jsonrpc_error(message, error),
};
match state
.registry
.cancel_async_job(&AsyncJobId::new(control.job_id.clone()), &now_rfc3339())
@@ -1588,7 +1677,7 @@ fn negotiate_post_response_mode(headers: &HeaderMap) -> Result<ResponseMode, Sta
}
}
if saw_json && saw_sse {
if saw_json || saw_sse {
return preferred.ok_or(StatusCode::NOT_ACCEPTABLE);
}
@@ -1620,6 +1709,21 @@ fn validate_get_accept_header(headers: &HeaderMap) -> Result<(), StatusCode> {
Err(StatusCode::NOT_ACCEPTABLE)
}
fn stream_session_belongs_to_tool(
session_record: &StreamSession,
tool: &PublishedAgentTool,
) -> bool {
session_record.workspace_id == tool.workspace_id
&& session_record.agent_id.as_ref() == Some(&tool.agent_id)
&& session_record.operation_id == tool.operation.id
}
fn async_job_belongs_to_tool(job: &AsyncJobHandle, tool: &PublishedAgentTool) -> bool {
job.workspace_id == tool.workspace_id
&& job.agent_id.as_ref() == Some(&tool.agent_id)
&& job.operation_id == tool.operation.id
}
fn protocol_version_from_headers(headers: &HeaderMap) -> Result<String, StatusCode> {
let Some(version) = headers.get(HEADER_MCP_PROTOCOL_VERSION) else {
return Ok(DEFAULT_PROTOCOL_VERSION.to_owned());