fix: harden streaming transport and ownership checks
This commit is contained in:
+107
-3
@@ -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());
|
||||
|
||||
Reference in New Issue
Block a user