1223 lines
39 KiB
Rust
1223 lines
39 KiB
Rust
use std::{
|
|
collections::BTreeMap,
|
|
convert::Infallible,
|
|
sync::Arc,
|
|
time::{Duration, Instant},
|
|
};
|
|
|
|
use axum::{
|
|
Json, Router,
|
|
extract::{Path, State},
|
|
http::{
|
|
HeaderMap, HeaderValue, StatusCode,
|
|
header::{AUTHORIZATION, RETRY_AFTER},
|
|
},
|
|
response::{IntoResponse, Response, sse::Event},
|
|
routing::get,
|
|
};
|
|
use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
|
|
use crank_core::{
|
|
AuthProfile, CoordinationStateStore, InvocationLevel, InvocationLog, InvocationLogId,
|
|
InvocationSource, InvocationStatus, OperationSecurityLevel, PlatformApiKeyScope, SecretId,
|
|
};
|
|
use crank_registry::{CreateInvocationLogRequest, PostgresRegistry, PublishedAgentTool};
|
|
use crank_runtime::{
|
|
RateLimitRejection, RequestRateLimiter, ResolvedAuth, RuntimeError, RuntimeExecutionRequest,
|
|
RuntimeExecutor, RuntimeOperation, RuntimeRequestContext, SecretCrypto,
|
|
};
|
|
use futures_util::stream;
|
|
use serde::{Deserialize, Serialize};
|
|
use serde_json::{Value, json};
|
|
use sha2::{Digest, Sha256};
|
|
use time::OffsetDateTime;
|
|
use tracing::info;
|
|
|
|
use crate::{
|
|
auth::{SharedMachineCredentialVerifier, VerifiedMachineCredential},
|
|
catalog::PublishedToolCatalog,
|
|
jsonrpc::{
|
|
DEFAULT_PROTOCOL_VERSION, is_notification, is_request, is_response, jsonrpc_error,
|
|
jsonrpc_result, method_name, negotiated_protocol_version, params, request_id,
|
|
},
|
|
manifest::tool_definitions,
|
|
session::{SessionState, SharedSessionStore},
|
|
tool_error::{
|
|
ToolErrorContract, generic_tool_error_contract, runtime_error_code,
|
|
tool_error_contract_from_runtime, tool_error_text, tool_error_value,
|
|
},
|
|
transport::{
|
|
AllowedOrigins, HEADER_MCP_SESSION_ID, ResponseMode, json_response,
|
|
negotiate_post_response_mode, protocol_version_from_headers, resolve_request_id,
|
|
session_id_from_headers, sse_response, transport_response, validate_get_accept_header,
|
|
validate_origin, validate_session_protocol_version, with_request_id_header,
|
|
},
|
|
};
|
|
|
|
const TRANSPORT_SESSION_TTL_MS: u64 = 86_400_000;
|
|
|
|
#[derive(Clone)]
|
|
pub struct AppState {
|
|
registry: PostgresRegistry,
|
|
catalog: PublishedToolCatalog,
|
|
runtime: RuntimeExecutor,
|
|
api_rate_limiter: RequestRateLimiter,
|
|
secret_crypto: SecretCrypto,
|
|
sessions: SharedSessionStore,
|
|
credential_verifier: SharedMachineCredentialVerifier,
|
|
allowed_origins: AllowedOrigins,
|
|
}
|
|
|
|
#[derive(Debug, Serialize, Deserialize)]
|
|
struct InitializeParams {
|
|
#[serde(rename = "protocolVersion")]
|
|
protocol_version: String,
|
|
}
|
|
|
|
#[derive(Debug, Serialize, Deserialize)]
|
|
struct ToolCallParams {
|
|
name: String,
|
|
#[serde(default)]
|
|
arguments: Value,
|
|
}
|
|
|
|
#[derive(Clone)]
|
|
struct ResolvedToolCall {
|
|
tool: PublishedAgentTool,
|
|
}
|
|
|
|
struct ToolCallExecution {
|
|
tool: PublishedAgentTool,
|
|
arguments: Value,
|
|
confirmation_token: Option<String>,
|
|
}
|
|
|
|
#[derive(Clone, Debug, Deserialize)]
|
|
struct AgentRoutePath {
|
|
workspace_slug: String,
|
|
agent_slug: String,
|
|
}
|
|
|
|
#[allow(clippy::too_many_arguments)]
|
|
pub fn build_app(
|
|
registry: PostgresRegistry,
|
|
refresh_interval: Duration,
|
|
public_base_url: Option<String>,
|
|
secret_crypto: SecretCrypto,
|
|
runtime: RuntimeExecutor,
|
|
api_rate_limiter: RequestRateLimiter,
|
|
coordination_store: Arc<dyn CoordinationStateStore>,
|
|
sessions: SharedSessionStore,
|
|
credential_verifier: SharedMachineCredentialVerifier,
|
|
) -> Router {
|
|
let state = Arc::new(AppState {
|
|
registry: registry.clone(),
|
|
catalog: PublishedToolCatalog::new(registry, refresh_interval, coordination_store),
|
|
runtime,
|
|
api_rate_limiter,
|
|
secret_crypto,
|
|
sessions,
|
|
credential_verifier,
|
|
allowed_origins: AllowedOrigins::new(public_base_url),
|
|
});
|
|
|
|
Router::new()
|
|
.route("/health", get(health))
|
|
.route(
|
|
"/v1/{workspace_slug}/{agent_slug}",
|
|
get(mcp_get).post(mcp_post).delete(mcp_delete),
|
|
)
|
|
.with_state(state)
|
|
}
|
|
|
|
async fn health() -> Json<Value> {
|
|
Json(json!({
|
|
"service": "mcp-server",
|
|
"status": "ok"
|
|
}))
|
|
}
|
|
|
|
async fn mcp_get(
|
|
Path(path): Path<AgentRoutePath>,
|
|
State(state): State<Arc<AppState>>,
|
|
headers: HeaderMap,
|
|
) -> Response {
|
|
if let Err(status) = validate_origin(&state.allowed_origins, &headers) {
|
|
return status.into_response();
|
|
}
|
|
|
|
if let Err(status) = validate_get_accept_header(&headers) {
|
|
return status.into_response();
|
|
}
|
|
|
|
if let Err(status) =
|
|
require_machine_access(&state, &path, &headers, PlatformApiKeyScope::Read).await
|
|
{
|
|
return status.into_response();
|
|
}
|
|
|
|
if let Err(rejection) = enforce_transport_rate_limit(&state, &path, &headers).await {
|
|
return rate_limited_status_response(rejection);
|
|
}
|
|
|
|
let session_id = match session_id_from_headers(&headers) {
|
|
Ok(Some(session_id)) => session_id,
|
|
Ok(None) => return StatusCode::BAD_REQUEST.into_response(),
|
|
Err(status) => return status.into_response(),
|
|
};
|
|
|
|
let Some(session) = (match state.sessions.get(&session_id).await {
|
|
Ok(session) => session,
|
|
Err(_) => return StatusCode::INTERNAL_SERVER_ERROR.into_response(),
|
|
}) else {
|
|
return StatusCode::NOT_FOUND.into_response();
|
|
};
|
|
|
|
if session.workspace_slug != path.workspace_slug || session.agent_slug != path.agent_slug {
|
|
return StatusCode::NOT_FOUND.into_response();
|
|
}
|
|
|
|
if !session.initialized {
|
|
return StatusCode::BAD_REQUEST.into_response();
|
|
}
|
|
|
|
if let Err(status) = validate_session_protocol_version(&headers, &session.protocol_version) {
|
|
return status.into_response();
|
|
}
|
|
|
|
sse_response(
|
|
StatusCode::OK,
|
|
stream::pending::<Result<Event, Infallible>>(),
|
|
Some(&session_id),
|
|
Some(&session.protocol_version),
|
|
)
|
|
}
|
|
|
|
async fn mcp_delete(
|
|
Path(path): Path<AgentRoutePath>,
|
|
State(state): State<Arc<AppState>>,
|
|
headers: HeaderMap,
|
|
) -> Response {
|
|
if let Err(status) =
|
|
require_machine_access(&state, &path, &headers, PlatformApiKeyScope::Read).await
|
|
{
|
|
return status.into_response();
|
|
}
|
|
|
|
if let Err(rejection) = enforce_transport_rate_limit(&state, &path, &headers).await {
|
|
return rate_limited_status_response(rejection);
|
|
}
|
|
|
|
match session_id_from_headers(&headers) {
|
|
Ok(Some(session_id)) => match state.sessions.get(&session_id).await {
|
|
Ok(Some(session))
|
|
if session.workspace_slug == path.workspace_slug
|
|
&& session.agent_slug == path.agent_slug =>
|
|
{
|
|
match state.sessions.delete(&session_id).await {
|
|
Ok(true) => StatusCode::NO_CONTENT.into_response(),
|
|
Ok(false) => StatusCode::NOT_FOUND.into_response(),
|
|
Err(_) => StatusCode::INTERNAL_SERVER_ERROR.into_response(),
|
|
}
|
|
}
|
|
Ok(Some(_)) => StatusCode::NOT_FOUND.into_response(),
|
|
Ok(None) => StatusCode::NOT_FOUND.into_response(),
|
|
Err(_) => StatusCode::INTERNAL_SERVER_ERROR.into_response(),
|
|
},
|
|
Ok(None) => StatusCode::BAD_REQUEST.into_response(),
|
|
Err(status) => status.into_response(),
|
|
}
|
|
}
|
|
|
|
async fn mcp_post(
|
|
Path(path): Path<AgentRoutePath>,
|
|
State(state): State<Arc<AppState>>,
|
|
headers: HeaderMap,
|
|
Json(message): Json<Value>,
|
|
) -> Response {
|
|
let transport_request_id = resolve_request_id(&headers);
|
|
info!(
|
|
request_id = %transport_request_id,
|
|
workspace_slug = %path.workspace_slug,
|
|
agent_slug = %path.agent_slug,
|
|
jsonrpc_method = method_name(&message).unwrap_or("<unknown>"),
|
|
"mcp request received"
|
|
);
|
|
|
|
if let Err(status) = validate_origin(&state.allowed_origins, &headers) {
|
|
return with_request_id_header(status.into_response(), &transport_request_id);
|
|
}
|
|
|
|
let response_mode = match negotiate_post_response_mode(&headers) {
|
|
Ok(mode) => mode,
|
|
Err(status) => {
|
|
return with_request_id_header(status.into_response(), &transport_request_id);
|
|
}
|
|
};
|
|
|
|
if is_response(&message) || is_notification(&message) && method_name(&message).is_none() {
|
|
return with_request_id_header(StatusCode::ACCEPTED.into_response(), &transport_request_id);
|
|
}
|
|
|
|
let protocol_version = match protocol_version_from_headers(&headers) {
|
|
Ok(value) => value,
|
|
Err(status) => {
|
|
return with_request_id_header(status.into_response(), &transport_request_id);
|
|
}
|
|
};
|
|
|
|
if let Err(rejection) = enforce_post_rate_limit(&state, &path, &headers).await {
|
|
return with_request_id_header(
|
|
rate_limited_jsonrpc_response(&message, response_mode, &protocol_version, rejection),
|
|
&transport_request_id,
|
|
);
|
|
}
|
|
|
|
if let Some(session_id) = headers.get(HEADER_MCP_SESSION_ID) {
|
|
if let Ok(session_id) = session_id.to_str() {
|
|
let session = match state.sessions.get(session_id).await {
|
|
Ok(session) => session,
|
|
Err(_) => {
|
|
return with_request_id_header(
|
|
StatusCode::INTERNAL_SERVER_ERROR.into_response(),
|
|
&transport_request_id,
|
|
);
|
|
}
|
|
};
|
|
if let Some(session) = session {
|
|
if let Err(status) =
|
|
validate_session_protocol_version(&headers, &session.protocol_version)
|
|
{
|
|
return with_request_id_header(status.into_response(), &transport_request_id);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
let required_scope = match method_name(&message) {
|
|
Some("tools/call") => PlatformApiKeyScope::Write,
|
|
_ => PlatformApiKeyScope::Read,
|
|
};
|
|
let credential = match require_machine_access(&state, &path, &headers, required_scope).await {
|
|
Ok(credential) => credential,
|
|
Err(status) => {
|
|
return with_request_id_header(status.into_response(), &transport_request_id);
|
|
}
|
|
};
|
|
|
|
let response = match method_name(&message) {
|
|
Some("initialize") if is_request(&message) => {
|
|
handle_initialize(state, &path, &message, response_mode).await
|
|
}
|
|
Some("notifications/initialized") if is_notification(&message) => {
|
|
handle_initialized_notification(state, &path, &headers).await
|
|
}
|
|
Some("ping") if is_request(&message) => {
|
|
let session = match require_initialized_session(&state, &path, &headers, &message).await
|
|
{
|
|
Ok(session) => session,
|
|
Err(response) => return response,
|
|
};
|
|
|
|
transport_response(
|
|
StatusCode::OK,
|
|
jsonrpc_result(
|
|
request_id(&message),
|
|
json!({ "protocolVersion": session.protocol_version }),
|
|
),
|
|
response_mode,
|
|
None,
|
|
Some(&session.protocol_version),
|
|
)
|
|
}
|
|
Some("tools/list") if is_request(&message) => {
|
|
let session = match require_initialized_session(&state, &path, &headers, &message).await
|
|
{
|
|
Ok(session) => session,
|
|
Err(response) => return response,
|
|
};
|
|
|
|
match state
|
|
.catalog
|
|
.list_tools(&session.workspace_slug, &session.agent_slug)
|
|
.await
|
|
{
|
|
Ok(tools) => {
|
|
let definitions = tools.iter().flat_map(tool_definitions).collect::<Vec<_>>();
|
|
|
|
transport_response(
|
|
StatusCode::OK,
|
|
jsonrpc_result(request_id(&message), json!({ "tools": definitions })),
|
|
response_mode,
|
|
None,
|
|
Some(&session.protocol_version),
|
|
)
|
|
}
|
|
Err(error) => internal_jsonrpc_error(&message, error),
|
|
}
|
|
}
|
|
Some("tools/call") if is_request(&message) => {
|
|
let session = match require_initialized_session(&state, &path, &headers, &message).await
|
|
{
|
|
Ok(session) => session,
|
|
Err(response) => return response,
|
|
};
|
|
let tool_call_params: ToolCallParams = match serde_json::from_value(params(&message)) {
|
|
Ok(value) => value,
|
|
Err(error) => {
|
|
return transport_response(
|
|
StatusCode::OK,
|
|
jsonrpc_error(request_id(&message), -32602, error.to_string()),
|
|
response_mode,
|
|
None,
|
|
Some(&session.protocol_version),
|
|
);
|
|
}
|
|
};
|
|
|
|
let mut arguments = if tool_call_params.arguments.is_null() {
|
|
json!({})
|
|
} else {
|
|
tool_call_params.arguments
|
|
};
|
|
let confirmation_token = take_confirmation_token(&mut arguments);
|
|
|
|
match state
|
|
.catalog
|
|
.list_tools(&session.workspace_slug, &session.agent_slug)
|
|
.await
|
|
{
|
|
Ok(tools) => match resolve_generated_tool(&tools, &tool_call_params.name) {
|
|
Some(resolved) => {
|
|
handle_tool_call(
|
|
state.clone(),
|
|
&session,
|
|
&message,
|
|
response_mode,
|
|
&credential,
|
|
resolved,
|
|
arguments,
|
|
confirmation_token,
|
|
&transport_request_id,
|
|
)
|
|
.await
|
|
}
|
|
None => transport_response(
|
|
StatusCode::OK,
|
|
jsonrpc_error(
|
|
request_id(&message),
|
|
-32602,
|
|
format!("tool {} was not found", tool_call_params.name),
|
|
),
|
|
response_mode,
|
|
None,
|
|
Some(&session.protocol_version),
|
|
),
|
|
},
|
|
Err(error) => internal_jsonrpc_error(&message, error),
|
|
}
|
|
}
|
|
Some(method) if is_notification(&message) => {
|
|
let _ = method;
|
|
StatusCode::ACCEPTED.into_response()
|
|
}
|
|
Some(method) => transport_response(
|
|
StatusCode::OK,
|
|
jsonrpc_error(
|
|
request_id(&message),
|
|
-32601,
|
|
format!("method {method} is not supported"),
|
|
),
|
|
response_mode,
|
|
None,
|
|
Some(&protocol_version),
|
|
),
|
|
None => transport_response(
|
|
StatusCode::BAD_REQUEST,
|
|
jsonrpc_error(Value::Null, -32600, "invalid JSON-RPC message"),
|
|
response_mode,
|
|
None,
|
|
Some(&protocol_version),
|
|
),
|
|
};
|
|
|
|
with_request_id_header(response, &transport_request_id)
|
|
}
|
|
|
|
#[allow(clippy::too_many_arguments)]
|
|
async fn handle_tool_call(
|
|
state: Arc<AppState>,
|
|
session: &SessionState,
|
|
message: &Value,
|
|
response_mode: ResponseMode,
|
|
credential: &VerifiedMachineCredential,
|
|
resolved: ResolvedToolCall,
|
|
arguments: Value,
|
|
confirmation_token: Option<String>,
|
|
transport_request_id: &str,
|
|
) -> Response {
|
|
if !credential_allows_security_level(credential, resolved.tool.operation.security_level) {
|
|
return tool_error_response(
|
|
message,
|
|
response_mode,
|
|
&session.protocol_version,
|
|
generic_tool_error_contract(
|
|
"machine_access_insufficient",
|
|
format!(
|
|
"machine access mode {} does not satisfy {} operation security",
|
|
serialize_machine_access_mode(credential.machine_access_mode),
|
|
serialize_security_level(resolved.tool.operation.security_level),
|
|
),
|
|
transport_request_id,
|
|
false,
|
|
Some("Используйте ключ агента с достаточным уровнем доступа."),
|
|
),
|
|
);
|
|
}
|
|
|
|
handle_base_tool_call(
|
|
state,
|
|
session,
|
|
message,
|
|
response_mode,
|
|
ToolCallExecution {
|
|
tool: resolved.tool,
|
|
arguments,
|
|
confirmation_token,
|
|
},
|
|
transport_request_id,
|
|
)
|
|
.await
|
|
}
|
|
|
|
async fn resolve_operation_auth(
|
|
state: &Arc<AppState>,
|
|
workspace_id: &crank_core::WorkspaceId,
|
|
execution_config: &crank_core::ExecutionConfig,
|
|
) -> Result<Option<ResolvedAuth>, RuntimeError> {
|
|
resolve_runtime_auth_for_task(
|
|
&state.registry,
|
|
&state.secret_crypto,
|
|
workspace_id,
|
|
execution_config,
|
|
)
|
|
.await
|
|
}
|
|
|
|
async fn resolve_runtime_auth_for_task(
|
|
registry: &PostgresRegistry,
|
|
secret_crypto: &SecretCrypto,
|
|
workspace_id: &crank_core::WorkspaceId,
|
|
execution_config: &crank_core::ExecutionConfig,
|
|
) -> Result<Option<ResolvedAuth>, RuntimeError> {
|
|
let Some(auth_profile_id) = execution_config.auth_profile_ref.as_ref() else {
|
|
return Ok(None);
|
|
};
|
|
|
|
let auth_profile = registry
|
|
.get_auth_profile(workspace_id, auth_profile_id)
|
|
.await
|
|
.map_err(|error| RuntimeError::SecretCrypto {
|
|
operation: "load auth profile",
|
|
details: error.to_string(),
|
|
})?
|
|
.ok_or_else(|| RuntimeError::MissingAuthProfile {
|
|
auth_profile_id: auth_profile_id.as_str().to_owned(),
|
|
})?;
|
|
|
|
resolve_auth_profile(registry, secret_crypto, workspace_id, &auth_profile)
|
|
.await
|
|
.map(Some)
|
|
}
|
|
|
|
async fn resolve_auth_profile(
|
|
registry: &PostgresRegistry,
|
|
secret_crypto: &SecretCrypto,
|
|
workspace_id: &crank_core::WorkspaceId,
|
|
auth_profile: &AuthProfile,
|
|
) -> Result<ResolvedAuth, RuntimeError> {
|
|
let mut secrets = BTreeMap::new();
|
|
let used_at = OffsetDateTime::now_utc();
|
|
|
|
for secret_id in auth_profile.config.secret_ids() {
|
|
let secret = registry
|
|
.get_secret(workspace_id, secret_id)
|
|
.await
|
|
.map_err(|error| RuntimeError::SecretCrypto {
|
|
operation: "load secret",
|
|
details: error.to_string(),
|
|
})?
|
|
.ok_or_else(|| RuntimeError::MissingSecret {
|
|
secret_id: secret_id.as_str().to_owned(),
|
|
})?;
|
|
let version = registry
|
|
.get_current_secret_version(workspace_id, secret_id)
|
|
.await
|
|
.map_err(|error| RuntimeError::SecretCrypto {
|
|
operation: "load current secret version",
|
|
details: error.to_string(),
|
|
})?
|
|
.ok_or_else(|| RuntimeError::MissingSecretVersion {
|
|
secret_id: secret_id.as_str().to_owned(),
|
|
version: secret.secret.current_version,
|
|
})?;
|
|
let plaintext = secret_crypto.decrypt(
|
|
&version.secret_version.key_version,
|
|
&version.secret_version.ciphertext,
|
|
)?;
|
|
registry
|
|
.touch_secret(workspace_id, secret_id, &used_at)
|
|
.await
|
|
.map_err(|error| RuntimeError::SecretCrypto {
|
|
operation: "touch secret",
|
|
details: error.to_string(),
|
|
})?;
|
|
secrets.insert(SecretId::new(secret_id.as_str()), plaintext);
|
|
}
|
|
|
|
ResolvedAuth::from_profile(auth_profile, &secrets)
|
|
}
|
|
|
|
async fn handle_base_tool_call(
|
|
state: Arc<AppState>,
|
|
session: &SessionState,
|
|
message: &Value,
|
|
response_mode: ResponseMode,
|
|
execution: ToolCallExecution,
|
|
transport_request_id: &str,
|
|
) -> Response {
|
|
let tool = execution.tool;
|
|
let arguments = execution.arguments;
|
|
let operation = runtime_operation(&tool);
|
|
let mut runtime_request_context = RuntimeRequestContext::from_request_id(transport_request_id)
|
|
.with_response_cache_scope(
|
|
tool.workspace_id.as_str().to_owned(),
|
|
tool.agent_id.as_str().to_owned(),
|
|
)
|
|
.with_metering_context(
|
|
tool.workspace_id.clone(),
|
|
Some(tool.agent_id.clone()),
|
|
InvocationSource::AgentToolCall,
|
|
);
|
|
if let Some(token) = execution.confirmation_token {
|
|
runtime_request_context = runtime_request_context.with_confirmation_token(token);
|
|
}
|
|
let request_preview = build_request_preview(&state.runtime, &operation, &arguments);
|
|
let started_at = Instant::now();
|
|
let resolved_auth =
|
|
resolve_operation_auth(&state, &tool.workspace_id, &operation.execution_config).await;
|
|
|
|
let result = match resolved_auth {
|
|
Ok(resolved_auth) => {
|
|
state
|
|
.runtime
|
|
.execute_request(
|
|
RuntimeExecutionRequest::new(&operation, &arguments)
|
|
.with_optional_auth(resolved_auth.as_ref())
|
|
.with_context(&runtime_request_context),
|
|
)
|
|
.await
|
|
}
|
|
Err(error) => Err(error),
|
|
};
|
|
|
|
match result {
|
|
Ok(output) => {
|
|
let _ = persist_invocation(
|
|
&state,
|
|
&tool,
|
|
InvocationRecord {
|
|
request_id: Some(transport_request_id),
|
|
tool_name: &tool.tool_name,
|
|
status: InvocationStatus::Ok,
|
|
level: InvocationLevel::Info,
|
|
message: "agent tool call completed",
|
|
status_code: None,
|
|
error_kind: None,
|
|
duration: started_at.elapsed(),
|
|
request_preview,
|
|
response_preview: output.clone(),
|
|
},
|
|
)
|
|
.await;
|
|
|
|
success_tool_response(message, response_mode, &session.protocol_version, output)
|
|
}
|
|
Err(error) => {
|
|
let _ = persist_invocation(
|
|
&state,
|
|
&tool,
|
|
InvocationRecord {
|
|
request_id: Some(transport_request_id),
|
|
tool_name: &tool.tool_name,
|
|
status: InvocationStatus::Error,
|
|
level: InvocationLevel::Error,
|
|
message: &error.to_string(),
|
|
status_code: None,
|
|
error_kind: Some(runtime_error_code(&error)),
|
|
duration: started_at.elapsed(),
|
|
request_preview,
|
|
response_preview: Value::Null,
|
|
},
|
|
)
|
|
.await;
|
|
|
|
tool_error_response(
|
|
message,
|
|
response_mode,
|
|
&session.protocol_version,
|
|
tool_error_contract_from_runtime(&error, transport_request_id),
|
|
)
|
|
}
|
|
}
|
|
}
|
|
|
|
async fn handle_initialize(
|
|
state: Arc<AppState>,
|
|
path: &AgentRoutePath,
|
|
message: &Value,
|
|
response_mode: ResponseMode,
|
|
) -> Response {
|
|
let initialize_params: InitializeParams = match serde_json::from_value(params(message)) {
|
|
Ok(value) => value,
|
|
Err(error) => {
|
|
return transport_response(
|
|
StatusCode::OK,
|
|
jsonrpc_error(request_id(message), -32602, error.to_string()),
|
|
response_mode,
|
|
None,
|
|
Some(DEFAULT_PROTOCOL_VERSION),
|
|
);
|
|
}
|
|
};
|
|
let Some(protocol_version) = negotiated_protocol_version(&initialize_params.protocol_version)
|
|
else {
|
|
return transport_response(
|
|
StatusCode::OK,
|
|
jsonrpc_error(
|
|
request_id(message),
|
|
-32602,
|
|
format!(
|
|
"unsupported protocol version {}",
|
|
initialize_params.protocol_version
|
|
),
|
|
),
|
|
response_mode,
|
|
None,
|
|
Some(DEFAULT_PROTOCOL_VERSION),
|
|
);
|
|
};
|
|
let now = OffsetDateTime::now_utc();
|
|
let expires_at = add_millis(now, TRANSPORT_SESSION_TTL_MS);
|
|
let session_id = match state
|
|
.sessions
|
|
.create(
|
|
protocol_version,
|
|
&path.workspace_slug,
|
|
&path.agent_slug,
|
|
now,
|
|
Some(expires_at),
|
|
)
|
|
.await
|
|
{
|
|
Ok(session_id) => session_id,
|
|
Err(error) => return internal_jsonrpc_error(message, error),
|
|
};
|
|
|
|
transport_response(
|
|
StatusCode::OK,
|
|
jsonrpc_result(
|
|
request_id(message),
|
|
json!({
|
|
"protocolVersion": protocol_version,
|
|
"capabilities": {
|
|
"tools": {
|
|
"listChanged": false
|
|
}
|
|
},
|
|
"serverInfo": {
|
|
"name": "crank-mcp-server",
|
|
"version": env!("CARGO_PKG_VERSION")
|
|
}
|
|
}),
|
|
),
|
|
response_mode,
|
|
Some(&session_id),
|
|
Some(protocol_version),
|
|
)
|
|
}
|
|
|
|
async fn handle_initialized_notification(
|
|
state: Arc<AppState>,
|
|
path: &AgentRoutePath,
|
|
headers: &HeaderMap,
|
|
) -> Response {
|
|
match session_id_from_headers(headers) {
|
|
Ok(Some(session_id)) => match state.sessions.get(&session_id).await {
|
|
Ok(Some(session))
|
|
if session.workspace_slug == path.workspace_slug
|
|
&& session.agent_slug == path.agent_slug =>
|
|
{
|
|
match state
|
|
.sessions
|
|
.mark_initialized(&session_id, OffsetDateTime::now_utc())
|
|
.await
|
|
{
|
|
Ok(true) => StatusCode::ACCEPTED.into_response(),
|
|
Ok(false) => StatusCode::NOT_FOUND.into_response(),
|
|
Err(_) => StatusCode::INTERNAL_SERVER_ERROR.into_response(),
|
|
}
|
|
}
|
|
Ok(Some(_)) => StatusCode::NOT_FOUND.into_response(),
|
|
Ok(None) => StatusCode::NOT_FOUND.into_response(),
|
|
Err(_) => StatusCode::INTERNAL_SERVER_ERROR.into_response(),
|
|
},
|
|
Ok(None) => StatusCode::BAD_REQUEST.into_response(),
|
|
Err(status) => status.into_response(),
|
|
}
|
|
}
|
|
|
|
async fn require_initialized_session(
|
|
state: &Arc<AppState>,
|
|
path: &AgentRoutePath,
|
|
headers: &HeaderMap,
|
|
message: &Value,
|
|
) -> Result<SessionState, Response> {
|
|
let session_id = match session_id_from_headers(headers) {
|
|
Ok(Some(session_id)) => session_id,
|
|
Ok(None) => return Err(StatusCode::BAD_REQUEST.into_response()),
|
|
Err(status) => return Err(status.into_response()),
|
|
};
|
|
|
|
let Some(session) = (match state.sessions.get(&session_id).await {
|
|
Ok(session) => session,
|
|
Err(_) => return Err(StatusCode::INTERNAL_SERVER_ERROR.into_response()),
|
|
}) else {
|
|
return Err(StatusCode::NOT_FOUND.into_response());
|
|
};
|
|
|
|
if session.workspace_slug != path.workspace_slug || session.agent_slug != path.agent_slug {
|
|
return Err(StatusCode::NOT_FOUND.into_response());
|
|
}
|
|
|
|
if !session.initialized {
|
|
return Err(json_response(
|
|
StatusCode::OK,
|
|
jsonrpc_error(request_id(message), -32002, "session is not initialized"),
|
|
None,
|
|
Some(&session.protocol_version),
|
|
));
|
|
}
|
|
|
|
Ok(session)
|
|
}
|
|
|
|
async fn require_machine_access(
|
|
state: &Arc<AppState>,
|
|
path: &AgentRoutePath,
|
|
headers: &HeaderMap,
|
|
required_scope: PlatformApiKeyScope,
|
|
) -> Result<VerifiedMachineCredential, StatusCode> {
|
|
let secret = bearer_token(headers).ok_or(StatusCode::UNAUTHORIZED)?;
|
|
let credential = resolve_machine_credential(state, path, secret).await?;
|
|
|
|
if !allows_scope(&credential.scopes, required_scope) {
|
|
return Err(StatusCode::FORBIDDEN);
|
|
}
|
|
|
|
Ok(credential)
|
|
}
|
|
|
|
async fn resolve_machine_credential(
|
|
state: &Arc<AppState>,
|
|
path: &AgentRoutePath,
|
|
token: &str,
|
|
) -> Result<VerifiedMachineCredential, StatusCode> {
|
|
if let Some(credential) = verify_static_agent_key(state, path, token).await? {
|
|
return Ok(credential);
|
|
}
|
|
|
|
state
|
|
.credential_verifier
|
|
.verify_bearer_token(&path.workspace_slug, &path.agent_slug, token)
|
|
.await
|
|
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?
|
|
.ok_or(StatusCode::UNAUTHORIZED)
|
|
}
|
|
|
|
async fn verify_static_agent_key(
|
|
state: &Arc<AppState>,
|
|
path: &AgentRoutePath,
|
|
secret: &str,
|
|
) -> Result<Option<VerifiedMachineCredential>, StatusCode> {
|
|
let secret_hash = hash_access_secret(secret);
|
|
let Some(api_key) = state
|
|
.registry
|
|
.get_platform_api_key_by_secret_for_agent_slug(
|
|
&path.workspace_slug,
|
|
&path.agent_slug,
|
|
&secret_hash,
|
|
)
|
|
.await
|
|
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?
|
|
else {
|
|
return Ok(None);
|
|
};
|
|
|
|
let used_at = OffsetDateTime::now_utc();
|
|
state
|
|
.registry
|
|
.touch_platform_api_key(&api_key.api_key.workspace_id, &api_key.api_key.id, &used_at)
|
|
.await
|
|
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
|
|
|
|
Ok(Some(VerifiedMachineCredential {
|
|
machine_access_mode: crank_core::MachineAccessMode::StaticAgentKey,
|
|
max_security_level: crank_core::OperationSecurityLevel::Standard,
|
|
scopes: api_key.api_key.scopes,
|
|
}))
|
|
}
|
|
|
|
fn bearer_token(headers: &HeaderMap) -> Option<&str> {
|
|
let value = headers.get(AUTHORIZATION)?.to_str().ok()?;
|
|
let (scheme, token) = value.split_once(' ')?;
|
|
if !scheme.eq_ignore_ascii_case("Bearer") || token.is_empty() {
|
|
return None;
|
|
}
|
|
|
|
Some(token)
|
|
}
|
|
|
|
async fn enforce_post_rate_limit(
|
|
state: &Arc<AppState>,
|
|
path: &AgentRoutePath,
|
|
headers: &HeaderMap,
|
|
) -> Result<(), RateLimitRejection> {
|
|
enforce_transport_rate_limit(state, path, headers).await
|
|
}
|
|
|
|
async fn enforce_transport_rate_limit(
|
|
state: &Arc<AppState>,
|
|
path: &AgentRoutePath,
|
|
headers: &HeaderMap,
|
|
) -> Result<(), RateLimitRejection> {
|
|
let key = rate_limit_key(path, headers);
|
|
state.api_rate_limiter.check(&key).await
|
|
}
|
|
|
|
fn rate_limit_key(path: &AgentRoutePath, headers: &HeaderMap) -> String {
|
|
if let Ok(Some(session_id)) = session_id_from_headers(headers) {
|
|
return format!(
|
|
"session:{}:{}:{}",
|
|
path.workspace_slug, path.agent_slug, session_id
|
|
);
|
|
}
|
|
|
|
if let Some(secret) = bearer_token(headers) {
|
|
return format!("api_key:{}", hash_access_secret(secret));
|
|
}
|
|
|
|
format!("workspace:{}:anonymous", path.workspace_slug)
|
|
}
|
|
|
|
fn rate_limited_jsonrpc_response(
|
|
message: &Value,
|
|
response_mode: ResponseMode,
|
|
protocol_version: &str,
|
|
rejection: RateLimitRejection,
|
|
) -> Response {
|
|
let payload = json!({
|
|
"jsonrpc": "2.0",
|
|
"id": request_id(message),
|
|
"error": {
|
|
"code": -32029,
|
|
"message": "request rate limit exceeded",
|
|
"data": {
|
|
"code": "request_rate_limited",
|
|
"retry_after_ms": rejection.retry_after_ms,
|
|
}
|
|
}
|
|
});
|
|
|
|
let mut response = transport_response(
|
|
StatusCode::TOO_MANY_REQUESTS,
|
|
payload,
|
|
response_mode,
|
|
None,
|
|
Some(protocol_version),
|
|
);
|
|
let retry_after_seconds = rejection.retry_after_ms.div_ceil(1000);
|
|
if let Ok(value) = HeaderValue::from_str(&retry_after_seconds.to_string()) {
|
|
response.headers_mut().insert(RETRY_AFTER, value);
|
|
}
|
|
response
|
|
}
|
|
|
|
fn rate_limited_status_response(rejection: RateLimitRejection) -> Response {
|
|
let mut response = StatusCode::TOO_MANY_REQUESTS.into_response();
|
|
let retry_after_seconds = rejection.retry_after_ms.div_ceil(1000);
|
|
if let Ok(value) = HeaderValue::from_str(&retry_after_seconds.to_string()) {
|
|
response.headers_mut().insert(RETRY_AFTER, value);
|
|
}
|
|
response
|
|
}
|
|
|
|
fn allows_scope(scopes: &[PlatformApiKeyScope], required_scope: PlatformApiKeyScope) -> bool {
|
|
match required_scope {
|
|
PlatformApiKeyScope::Read => scopes.iter().any(|scope| {
|
|
matches!(
|
|
scope,
|
|
PlatformApiKeyScope::Read
|
|
| PlatformApiKeyScope::Write
|
|
| PlatformApiKeyScope::Deploy
|
|
)
|
|
}),
|
|
PlatformApiKeyScope::Write => scopes.iter().any(|scope| {
|
|
matches!(
|
|
scope,
|
|
PlatformApiKeyScope::Write | PlatformApiKeyScope::Deploy
|
|
)
|
|
}),
|
|
PlatformApiKeyScope::Deploy => scopes
|
|
.iter()
|
|
.any(|scope| matches!(scope, PlatformApiKeyScope::Deploy)),
|
|
}
|
|
}
|
|
|
|
fn credential_allows_security_level(
|
|
credential: &VerifiedMachineCredential,
|
|
required_level: OperationSecurityLevel,
|
|
) -> bool {
|
|
security_level_rank(credential.max_security_level) >= security_level_rank(required_level)
|
|
}
|
|
|
|
fn security_level_rank(level: OperationSecurityLevel) -> u8 {
|
|
match level {
|
|
OperationSecurityLevel::Standard => 0,
|
|
}
|
|
}
|
|
|
|
fn serialize_security_level(level: OperationSecurityLevel) -> &'static str {
|
|
match level {
|
|
OperationSecurityLevel::Standard => "standard",
|
|
}
|
|
}
|
|
|
|
fn serialize_machine_access_mode(mode: crank_core::MachineAccessMode) -> &'static str {
|
|
match mode {
|
|
crank_core::MachineAccessMode::StaticAgentKey => "static_agent_key",
|
|
}
|
|
}
|
|
|
|
fn internal_jsonrpc_error(message: &Value, error: impl std::fmt::Display) -> Response {
|
|
transport_response(
|
|
StatusCode::INTERNAL_SERVER_ERROR,
|
|
jsonrpc_error(request_id(message), -32603, error.to_string()),
|
|
ResponseMode::Json,
|
|
None,
|
|
Some(DEFAULT_PROTOCOL_VERSION),
|
|
)
|
|
}
|
|
|
|
fn take_confirmation_token(arguments: &mut Value) -> Option<String> {
|
|
let Value::Object(object) = arguments else {
|
|
return None;
|
|
};
|
|
object
|
|
.remove("_crank_confirmation_token")
|
|
.and_then(|value| value.as_str().map(str::to_owned))
|
|
.filter(|value| !value.trim().is_empty())
|
|
}
|
|
|
|
fn build_request_preview(
|
|
runtime: &RuntimeExecutor,
|
|
operation: &RuntimeOperation,
|
|
arguments: &Value,
|
|
) -> Value {
|
|
match runtime.prepare_request(operation, arguments) {
|
|
Ok(prepared) => json!({
|
|
"path": prepared.path_params,
|
|
"query": prepared.query_params,
|
|
"headers": prepared.headers,
|
|
"body": prepared.body.unwrap_or(Value::Null)
|
|
}),
|
|
Err(_) => Value::Null,
|
|
}
|
|
}
|
|
|
|
struct InvocationRecord<'a> {
|
|
request_id: Option<&'a str>,
|
|
tool_name: &'a str,
|
|
status: InvocationStatus,
|
|
level: InvocationLevel,
|
|
message: &'a str,
|
|
status_code: Option<u16>,
|
|
error_kind: Option<&'a str>,
|
|
duration: Duration,
|
|
request_preview: Value,
|
|
response_preview: Value,
|
|
}
|
|
|
|
async fn persist_invocation(
|
|
state: &Arc<AppState>,
|
|
tool: &PublishedAgentTool,
|
|
record: InvocationRecord<'_>,
|
|
) -> Result<(), crank_registry::RegistryError> {
|
|
let created_at = OffsetDateTime::now_utc();
|
|
let duration_ms = u64::try_from(record.duration.as_millis()).unwrap_or(u64::MAX);
|
|
let log = InvocationLog {
|
|
id: InvocationLogId::new(format!("log_{}", uuid::Uuid::now_v7().simple())),
|
|
workspace_id: tool.workspace_id.clone(),
|
|
agent_id: Some(tool.agent_id.clone()),
|
|
operation_id: tool.operation.id.clone(),
|
|
source: InvocationSource::AgentToolCall,
|
|
level: record.level,
|
|
status: record.status,
|
|
tool_name: record.tool_name.to_owned(),
|
|
message: record.message.to_owned(),
|
|
request_id: record.request_id.map(ToOwned::to_owned),
|
|
status_code: record.status_code,
|
|
duration_ms,
|
|
error_kind: record.error_kind.map(ToOwned::to_owned),
|
|
request_preview: record.request_preview,
|
|
response_preview: record.response_preview,
|
|
created_at,
|
|
};
|
|
|
|
state
|
|
.registry
|
|
.create_invocation_log(CreateInvocationLogRequest { log: &log })
|
|
.await
|
|
}
|
|
|
|
fn success_tool_response(
|
|
message: &Value,
|
|
response_mode: ResponseMode,
|
|
protocol_version: &str,
|
|
output: Value,
|
|
) -> Response {
|
|
transport_response(
|
|
StatusCode::OK,
|
|
jsonrpc_result(
|
|
request_id(message),
|
|
json!({
|
|
"content": [
|
|
{
|
|
"type": "text",
|
|
"text": serde_json::to_string_pretty(&output).unwrap_or_else(|_| "{}".to_owned())
|
|
}
|
|
],
|
|
"structuredContent": output,
|
|
"isError": false
|
|
}),
|
|
),
|
|
response_mode,
|
|
None,
|
|
Some(protocol_version),
|
|
)
|
|
}
|
|
|
|
fn tool_error_response(
|
|
message: &Value,
|
|
response_mode: ResponseMode,
|
|
protocol_version: &str,
|
|
error: ToolErrorContract,
|
|
) -> Response {
|
|
let error_message = tool_error_text(&error);
|
|
let error_value = tool_error_value(&error);
|
|
|
|
transport_response(
|
|
StatusCode::OK,
|
|
jsonrpc_result(
|
|
request_id(message),
|
|
json!({
|
|
"content": [
|
|
{
|
|
"type": "text",
|
|
"text": error_message
|
|
}
|
|
],
|
|
"structuredContent": {
|
|
"error": error_value
|
|
},
|
|
"isError": true
|
|
}),
|
|
),
|
|
response_mode,
|
|
None,
|
|
Some(protocol_version),
|
|
)
|
|
}
|
|
|
|
fn add_millis(timestamp: OffsetDateTime, millis: u64) -> OffsetDateTime {
|
|
let delta = time::Duration::milliseconds(i64::try_from(millis).unwrap_or(i64::MAX));
|
|
|
|
timestamp + delta
|
|
}
|
|
|
|
fn hash_access_secret(secret: &str) -> String {
|
|
let digest = Sha256::digest(secret.as_bytes());
|
|
URL_SAFE_NO_PAD.encode(digest)
|
|
}
|
|
|
|
fn resolve_generated_tool(
|
|
tools: &[PublishedAgentTool],
|
|
tool_name: &str,
|
|
) -> Option<ResolvedToolCall> {
|
|
for tool in tools {
|
|
if tool.tool_name == tool_name {
|
|
return Some(ResolvedToolCall { tool: tool.clone() });
|
|
}
|
|
}
|
|
|
|
None
|
|
}
|
|
|
|
fn runtime_operation(tool: &PublishedAgentTool) -> RuntimeOperation {
|
|
let mut operation = RuntimeOperation::from(tool.operation.clone());
|
|
operation.tool_name = tool.tool_name.clone();
|
|
operation.tool_description.title = tool.tool_title.clone();
|
|
operation.tool_description.description = tool.tool_description.clone();
|
|
operation
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use axum::body::to_bytes;
|
|
use serde_json::{Value, json};
|
|
|
|
use super::{ResponseMode, tool_error_response};
|
|
use crate::jsonrpc::CURRENT_PROTOCOL_VERSION;
|
|
use crate::tool_error::generic_tool_error_contract;
|
|
|
|
#[tokio::test]
|
|
async fn tool_error_response_includes_structured_context() {
|
|
let response = tool_error_response(
|
|
&json!({"jsonrpc": "2.0", "id": "req-1"}),
|
|
ResponseMode::Json,
|
|
CURRENT_PROTOCOL_VERSION,
|
|
generic_tool_error_contract(
|
|
"streaming_payload_error",
|
|
"request root must be an object",
|
|
"req-1",
|
|
false,
|
|
Some("Проверьте параметры вызова инструмента."),
|
|
),
|
|
);
|
|
|
|
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
|
|
let payload: Value = serde_json::from_slice(&body).unwrap();
|
|
|
|
assert_eq!(
|
|
payload["result"]["structuredContent"]["error"],
|
|
json!({
|
|
"code": "streaming_payload_error",
|
|
"error_code": "streaming_payload_error",
|
|
"message": "request root must be an object",
|
|
"recoverable": false,
|
|
"request_id": "req-1",
|
|
"suggested_action": "Проверьте параметры вызова инструмента."
|
|
})
|
|
);
|
|
}
|
|
}
|