наблюдаемость: завершить базовый контур Community
Добавить структурированные журналы, метрики, трассировку и безопасный канал критических ошибок. Усилить границы рантайма, тесты, проверку зависимостей и сценарии развёртывания.
This commit is contained in:
@@ -3,18 +3,26 @@ name = "crank-adapter-rest"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
rust-version.workspace = true
|
||||
publish.workspace = true
|
||||
version.workspace = true
|
||||
|
||||
[dependencies]
|
||||
async-trait = "0.1"
|
||||
crank-core = { path = "../crank-core" }
|
||||
crank-trace = { path = "../crank-trace" }
|
||||
futures-util = "0.3"
|
||||
metrics.workspace = true
|
||||
opentelemetry.workspace = true
|
||||
reqwest = { workspace = true, features = ["stream"] }
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
thiserror.workspace = true
|
||||
tokio.workspace = true
|
||||
tracing.workspace = true
|
||||
tracing-opentelemetry.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
axum.workspace = true
|
||||
opentelemetry_sdk.workspace = true
|
||||
tokio.workspace = true
|
||||
tracing-subscriber.workspace = true
|
||||
|
||||
@@ -7,7 +7,9 @@ use std::{
|
||||
};
|
||||
|
||||
use crank_core::{HttpMethod, RestTarget};
|
||||
use crank_trace::{ErrorCategory, Stage, StageOutcome};
|
||||
use futures_util::StreamExt;
|
||||
use opentelemetry::{global, propagation::Injector, trace::TraceContextExt};
|
||||
use reqwest::{
|
||||
Client,
|
||||
dns::{Addrs, Name, Resolve, Resolving},
|
||||
@@ -15,6 +17,8 @@ use reqwest::{
|
||||
redirect,
|
||||
};
|
||||
use serde_json::Value;
|
||||
use tracing::{Instrument, Span};
|
||||
use tracing_opentelemetry::OpenTelemetrySpanExt;
|
||||
|
||||
use crate::{RestAdapterError, RestRequest, RestResponse};
|
||||
|
||||
@@ -66,10 +70,37 @@ impl RestAdapter {
|
||||
&self,
|
||||
target: &RestTarget,
|
||||
request: &RestRequest,
|
||||
) -> Result<RestResponse, RestAdapterError> {
|
||||
let started_at = std::time::Instant::now();
|
||||
let result = self.execute_inner(target, request).await;
|
||||
let outcome = match &result {
|
||||
Ok(_) => "success",
|
||||
Err(error) => upstream_outcome(error),
|
||||
};
|
||||
metrics::counter!(
|
||||
"crank_upstream_requests_total",
|
||||
"operation_kind" => "rest",
|
||||
"outcome" => outcome
|
||||
)
|
||||
.increment(1);
|
||||
metrics::histogram!(
|
||||
"crank_upstream_request_duration_seconds",
|
||||
"operation_kind" => "rest",
|
||||
"outcome" => outcome
|
||||
)
|
||||
.record(started_at.elapsed().as_secs_f64());
|
||||
result
|
||||
}
|
||||
|
||||
async fn execute_inner(
|
||||
&self,
|
||||
target: &RestTarget,
|
||||
request: &RestRequest,
|
||||
) -> Result<RestResponse, RestAdapterError> {
|
||||
let url = build_url(target, request)?;
|
||||
self.policy.validate_url(&url)?;
|
||||
let headers = build_headers(target, request)?;
|
||||
let mut headers = build_headers(target, request)?;
|
||||
apply_current_trace_context(&mut headers);
|
||||
let client =
|
||||
self.client
|
||||
.as_ref()
|
||||
@@ -85,23 +116,60 @@ impl RestAdapter {
|
||||
builder = builder.json(body);
|
||||
}
|
||||
|
||||
let response = builder.send().await?;
|
||||
let status = response.status();
|
||||
let headers = normalize_headers(response.headers());
|
||||
let body = decode_body(response, self.policy.max_response_bytes).await?;
|
||||
let upstream_span = Stage::UpstreamHttp.span();
|
||||
let result = async {
|
||||
let response = builder.send().await?;
|
||||
let status = response.status();
|
||||
let headers = normalize_headers(response.headers());
|
||||
let body = decode_body(response, self.policy.max_response_bytes).await?;
|
||||
|
||||
if !status.is_success() {
|
||||
return Err(RestAdapterError::UnexpectedStatus {
|
||||
status: status.as_u16(),
|
||||
if !status.is_success() {
|
||||
return Err(RestAdapterError::UnexpectedStatus {
|
||||
status: status.as_u16(),
|
||||
body,
|
||||
});
|
||||
}
|
||||
|
||||
Ok(RestResponse {
|
||||
status_code: status.as_u16(),
|
||||
headers,
|
||||
body,
|
||||
});
|
||||
})
|
||||
}
|
||||
.instrument(upstream_span.clone())
|
||||
.await;
|
||||
match &result {
|
||||
Ok(_) => StageOutcome::Success.record(&upstream_span),
|
||||
Err(_) => {
|
||||
StageOutcome::Error.record(&upstream_span);
|
||||
ErrorCategory::Upstream.record(&upstream_span);
|
||||
}
|
||||
}
|
||||
result
|
||||
}
|
||||
}
|
||||
|
||||
Ok(RestResponse {
|
||||
status_code: status.as_u16(),
|
||||
headers,
|
||||
body,
|
||||
})
|
||||
fn upstream_outcome(error: &RestAdapterError) -> &'static str {
|
||||
match error {
|
||||
RestAdapterError::UnexpectedStatus { status, .. } if (400..500).contains(status) => {
|
||||
"client_error"
|
||||
}
|
||||
RestAdapterError::UnexpectedStatus { status, .. } if (500..600).contains(status) => {
|
||||
"server_error"
|
||||
}
|
||||
RestAdapterError::UnexpectedStatus { .. } => "unexpected_status",
|
||||
RestAdapterError::Transport(error) if error.is_timeout() => "timeout",
|
||||
RestAdapterError::Transport(_) => "transport_error",
|
||||
RestAdapterError::ResponseTooLarge { .. } => "response_too_large",
|
||||
RestAdapterError::TargetNotAllowed { .. } => "rejected",
|
||||
RestAdapterError::WindowExpired => "window_expired",
|
||||
RestAdapterError::InvalidSseEvent => "invalid_response",
|
||||
RestAdapterError::InvalidBaseUrl { .. }
|
||||
| RestAdapterError::InvalidPathParameter { .. }
|
||||
| RestAdapterError::InvalidQueryParameter { .. }
|
||||
| RestAdapterError::InvalidHeaderName { .. }
|
||||
| RestAdapterError::InvalidHeaderValue { .. } => "invalid_request",
|
||||
RestAdapterError::InvalidConfiguration { .. } => "configuration",
|
||||
}
|
||||
}
|
||||
|
||||
@@ -407,6 +475,9 @@ fn insert_header(headers: &mut HeaderMap, name: &str, value: &str) -> Result<(),
|
||||
HeaderName::try_from(name).map_err(|_| RestAdapterError::InvalidHeaderName {
|
||||
header: name.to_owned(),
|
||||
})?;
|
||||
if is_trace_propagation_header(&header_name) {
|
||||
return Ok(());
|
||||
}
|
||||
let header_value =
|
||||
HeaderValue::try_from(value).map_err(|_| RestAdapterError::InvalidHeaderValue {
|
||||
header: name.to_owned(),
|
||||
@@ -416,6 +487,38 @@ fn insert_header(headers: &mut HeaderMap, name: &str, value: &str) -> Result<(),
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn is_trace_propagation_header(name: &HeaderName) -> bool {
|
||||
matches!(name.as_str(), "traceparent" | "tracestate" | "baggage")
|
||||
}
|
||||
|
||||
fn apply_current_trace_context(headers: &mut HeaderMap) {
|
||||
for header in ["traceparent", "tracestate", "baggage"] {
|
||||
headers.remove(header);
|
||||
}
|
||||
|
||||
let context = Span::current().context();
|
||||
if !context.span().span_context().is_valid() {
|
||||
return;
|
||||
}
|
||||
global::get_text_map_propagator(|propagator| {
|
||||
propagator.inject_context(&context, &mut ReqwestHeaderInjector(headers));
|
||||
});
|
||||
}
|
||||
|
||||
struct ReqwestHeaderInjector<'a>(&'a mut HeaderMap);
|
||||
|
||||
impl Injector for ReqwestHeaderInjector<'_> {
|
||||
fn set(&mut self, key: &str, value: String) {
|
||||
let Ok(name) = HeaderName::try_from(key) else {
|
||||
return;
|
||||
};
|
||||
let Ok(value) = HeaderValue::try_from(value) else {
|
||||
return;
|
||||
};
|
||||
self.0.insert(name, value);
|
||||
}
|
||||
}
|
||||
|
||||
async fn decode_body(
|
||||
response: reqwest::Response,
|
||||
max_response_bytes: usize,
|
||||
|
||||
@@ -26,13 +26,15 @@ impl ProtocolAdapter for RestAdapter {
|
||||
&self,
|
||||
target: &Target,
|
||||
prepared: &PreparedRequest,
|
||||
_context: &RuntimeRequestContext,
|
||||
context: &RuntimeRequestContext,
|
||||
) -> Result<AdapterResponse, ProtocolAdapterError> {
|
||||
let target = rest_target(target)?;
|
||||
let mut headers = prepared.headers.clone();
|
||||
headers.extend(context.outbound_headers());
|
||||
let request = RestRequest {
|
||||
path_params: prepared.path_params.clone(),
|
||||
query_params: prepared.query_params.clone(),
|
||||
headers: prepared.headers.clone(),
|
||||
headers,
|
||||
body: prepared.body.clone(),
|
||||
timeout_ms: prepared.timeout_ms,
|
||||
};
|
||||
|
||||
@@ -8,9 +8,19 @@ use axum::{
|
||||
routing::{get, post},
|
||||
};
|
||||
use crank_adapter_rest::{OutboundHttpPolicy, RestAdapter, RestAdapterError, RestRequest};
|
||||
use crank_core::{HttpMethod, RestTarget};
|
||||
use crank_core::{
|
||||
HttpMethod, PreparedRequest, ProtocolAdapter, RestTarget, RuntimeRequestContext, Target,
|
||||
};
|
||||
use opentelemetry::{
|
||||
global,
|
||||
trace::{TraceContextExt, TracerProvider as _},
|
||||
};
|
||||
use opentelemetry_sdk::{propagation::TraceContextPropagator, trace::SdkTracerProvider};
|
||||
use serde_json::{Value, json};
|
||||
use tokio::net::TcpListener;
|
||||
use tracing::Instrument;
|
||||
use tracing_opentelemetry::OpenTelemetrySpanExt;
|
||||
use tracing_subscriber::layer::SubscriberExt;
|
||||
|
||||
#[tokio::test]
|
||||
async fn executes_rest_request_and_normalizes_json_response() {
|
||||
@@ -45,6 +55,120 @@ async fn executes_rest_request_and_normalizes_json_response() {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn protocol_context_overrides_mapped_correlation_headers() {
|
||||
let base_url = spawn_test_server().await;
|
||||
let adapter = test_adapter();
|
||||
let target = Target::Rest(RestTarget {
|
||||
base_url,
|
||||
method: HttpMethod::Post,
|
||||
path_template: "/users/{user_id}".to_owned(),
|
||||
static_headers: BTreeMap::from([
|
||||
("x-request-id".to_owned(), "static-request".to_owned()),
|
||||
(
|
||||
"x-correlation-id".to_owned(),
|
||||
"static-correlation".to_owned(),
|
||||
),
|
||||
]),
|
||||
});
|
||||
let prepared = PreparedRequest {
|
||||
path_params: BTreeMap::from([("user_id".to_owned(), "42".to_owned())]),
|
||||
headers: BTreeMap::from([
|
||||
("x-request-id".to_owned(), "mapped-request".to_owned()),
|
||||
(
|
||||
"x-correlation-id".to_owned(),
|
||||
"mapped-correlation".to_owned(),
|
||||
),
|
||||
]),
|
||||
body: Some(json!({ "name": "Ada" })),
|
||||
timeout_ms: 1_000,
|
||||
..PreparedRequest::default()
|
||||
};
|
||||
let context = RuntimeRequestContext::new("req-runtime", "corr-runtime");
|
||||
|
||||
let response = adapter
|
||||
.invoke_unary(&target, &prepared, &context)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(response.body["request_id"], "req-runtime");
|
||||
assert_eq!(response.body["correlation_id"], "corr-runtime");
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "current_thread")]
|
||||
async fn current_trace_context_overrides_mapped_traceparent() {
|
||||
global::set_text_map_propagator(TraceContextPropagator::new());
|
||||
let provider = SdkTracerProvider::builder().build();
|
||||
let tracer = provider.tracer("rest-propagation-test");
|
||||
let subscriber =
|
||||
tracing_subscriber::registry().with(tracing_opentelemetry::layer().with_tracer(tracer));
|
||||
let dispatch = tracing::Dispatch::new(subscriber);
|
||||
let _dispatch_guard = tracing::dispatcher::set_default(&dispatch);
|
||||
let span = tracing::info_span!("runtime.execute");
|
||||
let context = span.context();
|
||||
let expected_trace_id = context.span().span_context().trace_id().to_string();
|
||||
let base_url = spawn_test_server().await;
|
||||
let target = RestTarget {
|
||||
base_url,
|
||||
method: HttpMethod::Post,
|
||||
path_template: "/users/{user_id}".to_owned(),
|
||||
static_headers: BTreeMap::new(),
|
||||
};
|
||||
let request = RestRequest {
|
||||
path_params: BTreeMap::from([("user_id".to_owned(), "42".to_owned())]),
|
||||
query_params: BTreeMap::new(),
|
||||
headers: BTreeMap::from([(
|
||||
"traceparent".to_owned(),
|
||||
"00-11111111111111111111111111111111-2222222222222222-01".to_owned(),
|
||||
)]),
|
||||
body: Some(json!({ "name": "Ada" })),
|
||||
timeout_ms: 1_000,
|
||||
};
|
||||
|
||||
let response = test_adapter()
|
||||
.execute(&target, &request)
|
||||
.instrument(span)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
&response.body["traceparent"].as_str().unwrap()[3..35],
|
||||
expected_trace_id
|
||||
);
|
||||
provider.shutdown().unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn user_configured_propagation_headers_are_removed_without_trusted_context() {
|
||||
let base_url = spawn_test_server().await;
|
||||
let target = RestTarget {
|
||||
base_url,
|
||||
method: HttpMethod::Post,
|
||||
path_template: "/users/{user_id}".to_owned(),
|
||||
static_headers: BTreeMap::from([
|
||||
(
|
||||
"traceparent".to_owned(),
|
||||
"untrusted\ninvalid-value".to_owned(),
|
||||
),
|
||||
("tracestate".to_owned(), "vendor=value".to_owned()),
|
||||
("baggage".to_owned(), "secret=must-not-leave".to_owned()),
|
||||
]),
|
||||
};
|
||||
let request = RestRequest {
|
||||
path_params: BTreeMap::from([("user_id".to_owned(), "42".to_owned())]),
|
||||
query_params: BTreeMap::new(),
|
||||
headers: BTreeMap::new(),
|
||||
body: Some(json!({ "name": "Ada" })),
|
||||
timeout_ms: 1_000,
|
||||
};
|
||||
|
||||
let response = test_adapter().execute(&target, &request).await.unwrap();
|
||||
|
||||
assert!(response.body.get("traceparent").is_none());
|
||||
assert!(response.body.get("tracestate").is_none());
|
||||
assert!(response.body.get("baggage").is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn returns_unexpected_status_with_normalized_body() {
|
||||
let base_url = spawn_test_server().await;
|
||||
@@ -186,14 +310,57 @@ async fn create_user(
|
||||
.get("x-static")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.unwrap_or_default();
|
||||
let request_id = headers
|
||||
.get("x-request-id")
|
||||
.and_then(|value| value.to_str().ok());
|
||||
let correlation_id = headers
|
||||
.get("x-correlation-id")
|
||||
.and_then(|value| value.to_str().ok());
|
||||
let traceparent = headers
|
||||
.get("traceparent")
|
||||
.and_then(|value| value.to_str().ok());
|
||||
let tracestate = headers
|
||||
.get("tracestate")
|
||||
.and_then(|value| value.to_str().ok());
|
||||
let baggage = headers.get("baggage").and_then(|value| value.to_str().ok());
|
||||
|
||||
Json(json!({
|
||||
let mut response = json!({
|
||||
"id": user_id,
|
||||
"query": query.get("expand").cloned().unwrap_or_default(),
|
||||
"trace": trace,
|
||||
"static": static_header,
|
||||
"payload": payload
|
||||
}))
|
||||
});
|
||||
let response = response.as_object_mut().unwrap();
|
||||
if let Some(request_id) = request_id {
|
||||
response.insert(
|
||||
"request_id".to_owned(),
|
||||
Value::String(request_id.to_owned()),
|
||||
);
|
||||
}
|
||||
if let Some(correlation_id) = correlation_id {
|
||||
response.insert(
|
||||
"correlation_id".to_owned(),
|
||||
Value::String(correlation_id.to_owned()),
|
||||
);
|
||||
}
|
||||
if let Some(traceparent) = traceparent {
|
||||
response.insert(
|
||||
"traceparent".to_owned(),
|
||||
Value::String(traceparent.to_owned()),
|
||||
);
|
||||
}
|
||||
if let Some(tracestate) = tracestate {
|
||||
response.insert(
|
||||
"tracestate".to_owned(),
|
||||
Value::String(tracestate.to_owned()),
|
||||
);
|
||||
}
|
||||
if let Some(baggage) = baggage {
|
||||
response.insert("baggage".to_owned(), Value::String(baggage.to_owned()));
|
||||
}
|
||||
|
||||
Json(Value::Object(response.clone()))
|
||||
}
|
||||
|
||||
async fn fail() -> (axum::http::StatusCode, Json<Value>) {
|
||||
|
||||
@@ -3,6 +3,7 @@ name = "crank-community-auth"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
rust-version.workspace = true
|
||||
publish.workspace = true
|
||||
version.workspace = true
|
||||
|
||||
[dependencies]
|
||||
|
||||
@@ -52,7 +52,11 @@ impl IdentityProvider for PasswordIdentityProvider {
|
||||
&self.password_pepper,
|
||||
&user.password_hash,
|
||||
) {
|
||||
debug!(email = %payload.email, "password identity provider rejected credentials");
|
||||
debug!(
|
||||
name: "auth.password.rejected",
|
||||
identity_provider = "password",
|
||||
"password identity provider rejected credentials"
|
||||
);
|
||||
return Err(IdentityError::BadCredentials);
|
||||
}
|
||||
|
||||
|
||||
@@ -3,6 +3,7 @@ name = "crank-community-mcp"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
rust-version.workspace = true
|
||||
publish.workspace = true
|
||||
version.workspace = true
|
||||
|
||||
[dependencies]
|
||||
@@ -11,10 +12,13 @@ axum.workspace = true
|
||||
base64.workspace = true
|
||||
crank-adapter-rest = { path = "../crank-adapter-rest" }
|
||||
crank-core = { path = "../crank-core" }
|
||||
crank-observability = { path = "../crank-observability" }
|
||||
crank-registry = { path = "../crank-registry" }
|
||||
crank-runtime = { path = "../crank-runtime" }
|
||||
crank-schema = { path = "../crank-schema" }
|
||||
crank-trace = { path = "../crank-trace" }
|
||||
futures-util = "0.3"
|
||||
metrics.workspace = true
|
||||
reqwest.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
@@ -29,3 +33,7 @@ uuid.workspace = true
|
||||
[dev-dependencies]
|
||||
crank-mapping = { path = "../crank-mapping" }
|
||||
crank-test-support = { path = "../crank-test-support" }
|
||||
opentelemetry.workspace = true
|
||||
opentelemetry_sdk.workspace = true
|
||||
tracing-opentelemetry.workspace = true
|
||||
tracing-subscriber.workspace = true
|
||||
|
||||
@@ -1,27 +1,54 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use axum::http::{HeaderMap, StatusCode, header::AUTHORIZATION};
|
||||
use axum::{
|
||||
http::{HeaderMap, StatusCode, header::AUTHORIZATION},
|
||||
response::{IntoResponse, Response},
|
||||
};
|
||||
use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
|
||||
use crank_core::{OperationSecurityLevel, PlatformApiKeyScope};
|
||||
use crank_trace::{DbOperation, ErrorCategory, Stage, StageOutcome, observe_db_query};
|
||||
use sha2::{Digest, Sha256};
|
||||
use time::OffsetDateTime;
|
||||
use tracing::Instrument;
|
||||
|
||||
use crate::{
|
||||
app::{AgentRoutePath, AppState},
|
||||
auth::VerifiedMachineCredential,
|
||||
};
|
||||
|
||||
#[derive(Clone, Copy, Debug)]
|
||||
pub(super) enum MachineAccessError {
|
||||
Denied(StatusCode),
|
||||
Unavailable,
|
||||
}
|
||||
|
||||
impl MachineAccessError {
|
||||
pub(super) fn is_denied(self) -> bool {
|
||||
matches!(self, Self::Denied(_))
|
||||
}
|
||||
}
|
||||
|
||||
impl IntoResponse for MachineAccessError {
|
||||
fn into_response(self) -> Response {
|
||||
match self {
|
||||
Self::Denied(status) => status.into_response(),
|
||||
Self::Unavailable => StatusCode::INTERNAL_SERVER_ERROR.into_response(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) 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)?;
|
||||
) -> Result<VerifiedMachineCredential, MachineAccessError> {
|
||||
let secret =
|
||||
bearer_token(headers).ok_or(MachineAccessError::Denied(StatusCode::UNAUTHORIZED))?;
|
||||
let credential = resolve_machine_credential(state, path, secret).await?;
|
||||
|
||||
if !allows_scope(&credential.scopes, required_scope) {
|
||||
return Err(StatusCode::FORBIDDEN);
|
||||
return Err(MachineAccessError::Denied(StatusCode::FORBIDDEN));
|
||||
}
|
||||
|
||||
Ok(credential)
|
||||
@@ -35,15 +62,18 @@ pub(super) async fn require_approval_access(
|
||||
) -> Result<crank_registry::PlatformApiKeyRecord, StatusCode> {
|
||||
let secret = bearer_token(headers).ok_or(StatusCode::UNAUTHORIZED)?;
|
||||
let secret_hash = hash_access_secret(secret);
|
||||
let Some(api_key) = state
|
||||
.registry
|
||||
.get_approval_api_key_by_secret_for_agent_slug(
|
||||
&path.workspace_slug,
|
||||
&path.agent_slug,
|
||||
&secret_hash,
|
||||
)
|
||||
.await
|
||||
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?
|
||||
let Some(api_key) = observe_db_query(
|
||||
DbOperation::MachineAccessRead,
|
||||
state
|
||||
.registry
|
||||
.get_approval_api_key_by_secret_for_agent_slug(
|
||||
&path.workspace_slug,
|
||||
&path.agent_slug,
|
||||
&secret_hash,
|
||||
),
|
||||
)
|
||||
.await
|
||||
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?
|
||||
else {
|
||||
return Err(StatusCode::UNAUTHORIZED);
|
||||
};
|
||||
@@ -53,11 +83,16 @@ pub(super) async fn require_approval_access(
|
||||
}
|
||||
|
||||
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)?;
|
||||
observe_db_query(
|
||||
DbOperation::MachineAccessTouch,
|
||||
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(api_key)
|
||||
}
|
||||
@@ -100,7 +135,7 @@ async fn resolve_machine_credential(
|
||||
state: &Arc<AppState>,
|
||||
path: &AgentRoutePath,
|
||||
token: &str,
|
||||
) -> Result<VerifiedMachineCredential, StatusCode> {
|
||||
) -> Result<VerifiedMachineCredential, MachineAccessError> {
|
||||
if let Some(credential) = verify_static_agent_key(state, path, token).await? {
|
||||
return Ok(credential);
|
||||
}
|
||||
@@ -109,35 +144,68 @@ async fn resolve_machine_credential(
|
||||
.credential_verifier
|
||||
.verify_bearer_token(&path.workspace_slug, &path.agent_slug, token)
|
||||
.await
|
||||
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?
|
||||
.ok_or(StatusCode::UNAUTHORIZED)
|
||||
.map_err(|_| MachineAccessError::Unavailable)?
|
||||
.ok_or(MachineAccessError::Denied(StatusCode::UNAUTHORIZED))
|
||||
}
|
||||
|
||||
async fn verify_static_agent_key(
|
||||
state: &Arc<AppState>,
|
||||
path: &AgentRoutePath,
|
||||
secret: &str,
|
||||
) -> Result<Option<VerifiedMachineCredential>, StatusCode> {
|
||||
) -> Result<Option<VerifiedMachineCredential>, MachineAccessError> {
|
||||
let secret_hash = hash_access_secret(secret);
|
||||
let Some(api_key) = state
|
||||
let read_span = Stage::DbQuery
|
||||
.db_span(DbOperation::MachineAccessRead)
|
||||
.expect("database stage");
|
||||
let api_key_result = state
|
||||
.registry
|
||||
.get_platform_api_key_by_secret_for_agent_slug(
|
||||
&path.workspace_slug,
|
||||
&path.agent_slug,
|
||||
&secret_hash,
|
||||
)
|
||||
.instrument(read_span.clone())
|
||||
.await
|
||||
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?
|
||||
else {
|
||||
.map_err(|_| MachineAccessError::Unavailable);
|
||||
let api_key = match api_key_result {
|
||||
Ok(api_key) => {
|
||||
StageOutcome::Success.record(&read_span);
|
||||
drop(read_span);
|
||||
api_key
|
||||
}
|
||||
Err(status) => {
|
||||
StageOutcome::Error.record(&read_span);
|
||||
ErrorCategory::Database.record(&read_span);
|
||||
drop(read_span);
|
||||
return Err(status);
|
||||
}
|
||||
};
|
||||
let Some(api_key) = api_key else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let used_at = OffsetDateTime::now_utc();
|
||||
state
|
||||
let touch_span = Stage::DbQuery
|
||||
.db_span(DbOperation::MachineAccessTouch)
|
||||
.expect("database stage");
|
||||
let touch_result = state
|
||||
.registry
|
||||
.touch_platform_api_key(&api_key.api_key.workspace_id, &api_key.api_key.id, &used_at)
|
||||
.instrument(touch_span.clone())
|
||||
.await
|
||||
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
|
||||
.map_err(|_| MachineAccessError::Unavailable);
|
||||
match touch_result {
|
||||
Ok(()) => {
|
||||
StageOutcome::Success.record(&touch_span);
|
||||
drop(touch_span);
|
||||
}
|
||||
Err(status) => {
|
||||
StageOutcome::Error.record(&touch_span);
|
||||
ErrorCategory::Database.record(&touch_span);
|
||||
drop(touch_span);
|
||||
return Err(status);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(Some(VerifiedMachineCredential {
|
||||
machine_access_mode: crank_core::MachineAccessMode::StaticAgentKey,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,117 @@
|
||||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use crank_core::{
|
||||
InvocationLevel, InvocationLog, InvocationLogId, InvocationSource, InvocationStatus,
|
||||
};
|
||||
use crank_registry::{
|
||||
CreateInvocationLogRequest, InvocationHistoryWriteOutcome, PublishedAgentTool,
|
||||
};
|
||||
use crank_trace::{DbOperation, ErrorCategory, Stage, StageOutcome};
|
||||
use serde_json::Value;
|
||||
use time::OffsetDateTime;
|
||||
use tracing::{Instrument, warn};
|
||||
|
||||
use super::AppState;
|
||||
|
||||
pub(crate) struct InvocationRecord<'a> {
|
||||
pub(crate) request_id: Option<&'a str>,
|
||||
pub(crate) tool_name: &'a str,
|
||||
pub(crate) status: InvocationStatus,
|
||||
pub(crate) level: InvocationLevel,
|
||||
pub(crate) message: &'a str,
|
||||
pub(crate) status_code: Option<u16>,
|
||||
pub(crate) error_kind: Option<&'a str>,
|
||||
pub(crate) duration: Duration,
|
||||
pub(crate) request_preview: Value,
|
||||
pub(crate) response_preview: Value,
|
||||
}
|
||||
|
||||
pub(crate) async fn persist_invocation(
|
||||
state: &Arc<AppState>,
|
||||
tool: &PublishedAgentTool,
|
||||
record: InvocationRecord<'_>,
|
||||
) -> InvocationHistoryWriteOutcome {
|
||||
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: u64::try_from(record.duration.as_millis()).unwrap_or(u64::MAX),
|
||||
error_kind: record.error_kind.map(ToOwned::to_owned),
|
||||
request_preview: record.request_preview,
|
||||
response_preview: record.response_preview,
|
||||
created_at: OffsetDateTime::now_utc(),
|
||||
};
|
||||
|
||||
let history_span = Stage::HistoryWrite.span();
|
||||
let (outcome, db_span) = async {
|
||||
let db_span = Stage::DbQuery
|
||||
.db_span(DbOperation::InvocationHistoryWrite)
|
||||
.expect("database stage");
|
||||
let outcome = state
|
||||
.registry
|
||||
.create_invocation_log(CreateInvocationLogRequest { log: &log })
|
||||
.instrument(db_span.clone())
|
||||
.await;
|
||||
(outcome, db_span)
|
||||
}
|
||||
.instrument(history_span.clone())
|
||||
.await;
|
||||
match outcome {
|
||||
InvocationHistoryWriteOutcome::Recorded => {
|
||||
StageOutcome::Success.record(&db_span);
|
||||
StageOutcome::Success.record(&history_span);
|
||||
}
|
||||
InvocationHistoryWriteOutcome::Lost(_) => {
|
||||
StageOutcome::Error.record(&db_span);
|
||||
ErrorCategory::Database.record(&db_span);
|
||||
StageOutcome::Error.record(&history_span);
|
||||
ErrorCategory::History.record(&history_span);
|
||||
}
|
||||
}
|
||||
drop(db_span);
|
||||
drop(history_span);
|
||||
observe_invocation_history_outcome(
|
||||
outcome,
|
||||
record.request_id,
|
||||
record.status,
|
||||
"agent_tool_call",
|
||||
);
|
||||
outcome
|
||||
}
|
||||
|
||||
pub(super) fn observe_invocation_history_outcome(
|
||||
outcome: InvocationHistoryWriteOutcome,
|
||||
request_id: Option<&str>,
|
||||
status: InvocationStatus,
|
||||
source: &'static str,
|
||||
) {
|
||||
let Some(loss) = outcome.loss() else {
|
||||
return;
|
||||
};
|
||||
crank_observability::record_operational_incident(
|
||||
crank_observability::OperationalIncident::InvocationHistoryLost,
|
||||
);
|
||||
warn!(
|
||||
name: "mcp.invocation_history.lost",
|
||||
request_id = request_id.unwrap_or_default(),
|
||||
source,
|
||||
invocation_status = invocation_status_label(status),
|
||||
error_category = loss.category.as_str(),
|
||||
"invocation history was not recorded"
|
||||
);
|
||||
}
|
||||
|
||||
fn invocation_status_label(status: InvocationStatus) -> &'static str {
|
||||
match status {
|
||||
InvocationStatus::Ok => "ok",
|
||||
InvocationStatus::Error => "error",
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,93 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use axum::http::StatusCode;
|
||||
use serde_json::Value;
|
||||
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
|
||||
|
||||
use crate::{
|
||||
jsonrpc::{is_notification, is_response, method_name},
|
||||
transport::ResponseMode,
|
||||
};
|
||||
|
||||
pub(super) struct McpRequestMetrics {
|
||||
method: &'static str,
|
||||
response_mode: &'static str,
|
||||
outcome: &'static str,
|
||||
}
|
||||
|
||||
impl McpRequestMetrics {
|
||||
pub(super) fn new(message: &Value) -> Self {
|
||||
Self {
|
||||
method: normalized_mcp_method(message),
|
||||
response_mode: "unknown",
|
||||
outcome: "rejected",
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn set_response_mode(&mut self, mode: ResponseMode) -> ResponseMode {
|
||||
self.response_mode = match mode {
|
||||
ResponseMode::Json => "json",
|
||||
ResponseMode::Sse => "sse",
|
||||
};
|
||||
mode
|
||||
}
|
||||
|
||||
pub(super) fn complete(&mut self, status: StatusCode) {
|
||||
self.outcome = match status.as_u16() {
|
||||
200..=299 => "success",
|
||||
400..=499 => "client_error",
|
||||
500..=599 => "server_error",
|
||||
_ => "other",
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for McpRequestMetrics {
|
||||
fn drop(&mut self) {
|
||||
::metrics::counter!(
|
||||
"crank_mcp_requests_total",
|
||||
"method" => self.method,
|
||||
"response_mode" => self.response_mode,
|
||||
"outcome" => self.outcome
|
||||
)
|
||||
.increment(1);
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn normalized_mcp_method(message: &Value) -> &'static str {
|
||||
match method_name(message) {
|
||||
Some("initialize") => "initialize",
|
||||
Some("notifications/initialized") => "initialized",
|
||||
Some("ping") => "ping",
|
||||
Some("tools/list") => "tools_list",
|
||||
Some("tools/call") => "tools_call",
|
||||
Some(_) if is_notification(message) => "notification",
|
||||
Some(_) => "unsupported",
|
||||
None if is_response(message) => "response",
|
||||
None => "invalid",
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) struct ActiveSessionGuard {
|
||||
_permit: OwnedSemaphorePermit,
|
||||
}
|
||||
|
||||
impl ActiveSessionGuard {
|
||||
pub(super) fn try_acquire(slots: &Arc<Semaphore>) -> Result<Self, ()> {
|
||||
let permit = Arc::clone(slots).try_acquire_owned().map_err(|_| {
|
||||
::metrics::counter!(
|
||||
"crank_runtime_limit_rejections_total",
|
||||
"stage" => "mcp_session"
|
||||
)
|
||||
.increment(1);
|
||||
})?;
|
||||
::metrics::gauge!("crank_mcp_active_sessions").increment(1.0);
|
||||
Ok(Self { _permit: permit })
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for ActiveSessionGuard {
|
||||
fn drop(&mut self) {
|
||||
::metrics::gauge!("crank_mcp_active_sessions").decrement(1.0);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use axum::http::{HeaderMap, StatusCode};
|
||||
use crank_core::PlatformApiKeyScope;
|
||||
use crank_registry::PlatformApiKeyRecord;
|
||||
use crank_runtime::RateLimitCheckError;
|
||||
use crank_trace::{ErrorCategory, Stage, StageOutcome};
|
||||
use tracing::Instrument;
|
||||
|
||||
use crate::{
|
||||
access::{MachineAccessError, require_approval_access, require_machine_access},
|
||||
app::{AgentRoutePath, AppState},
|
||||
auth::VerifiedMachineCredential,
|
||||
rate_limit::enforce_transport_rate_limit,
|
||||
};
|
||||
|
||||
pub(super) async fn enforce_traced_rate_limit(
|
||||
state: &Arc<AppState>,
|
||||
path: &AgentRoutePath,
|
||||
headers: &HeaderMap,
|
||||
) -> Result<(), RateLimitCheckError> {
|
||||
let span = Stage::McpRateLimit.span();
|
||||
let result = enforce_transport_rate_limit(state, path, headers)
|
||||
.instrument(span.clone())
|
||||
.await;
|
||||
match &result {
|
||||
Ok(()) => StageOutcome::Allowed.record(&span),
|
||||
Err(RateLimitCheckError::Rejected(_)) => {
|
||||
StageOutcome::Denied.record(&span);
|
||||
ErrorCategory::RateLimit.record(&span);
|
||||
}
|
||||
Err(RateLimitCheckError::StoreUnavailable) => {
|
||||
StageOutcome::Error.record(&span);
|
||||
ErrorCategory::Internal.record(&span);
|
||||
}
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
pub(super) async fn require_traced_machine_access(
|
||||
state: &Arc<AppState>,
|
||||
path: &AgentRoutePath,
|
||||
headers: &HeaderMap,
|
||||
required_scope: PlatformApiKeyScope,
|
||||
) -> Result<VerifiedMachineCredential, MachineAccessError> {
|
||||
let span = Stage::McpAccessCheck.span();
|
||||
let result = require_machine_access(state, path, headers, required_scope)
|
||||
.instrument(span.clone())
|
||||
.await;
|
||||
match &result {
|
||||
Ok(_) => StageOutcome::Allowed.record(&span),
|
||||
Err(error) if error.is_denied() => {
|
||||
StageOutcome::Denied.record(&span);
|
||||
ErrorCategory::Access.record(&span);
|
||||
}
|
||||
Err(_) => {
|
||||
StageOutcome::Error.record(&span);
|
||||
ErrorCategory::Internal.record(&span);
|
||||
}
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
pub(super) async fn require_traced_approval_access(
|
||||
state: &Arc<AppState>,
|
||||
path: &AgentRoutePath,
|
||||
headers: &HeaderMap,
|
||||
required_scope: PlatformApiKeyScope,
|
||||
) -> Result<PlatformApiKeyRecord, StatusCode> {
|
||||
let span = Stage::McpAccessCheck.span();
|
||||
let result = require_approval_access(state, path, headers, required_scope)
|
||||
.instrument(span.clone())
|
||||
.await;
|
||||
match &result {
|
||||
Ok(_) => StageOutcome::Allowed.record(&span),
|
||||
Err(status) if *status == StatusCode::UNAUTHORIZED || *status == StatusCode::FORBIDDEN => {
|
||||
StageOutcome::Denied.record(&span);
|
||||
ErrorCategory::Access.record(&span);
|
||||
}
|
||||
Err(_) => {
|
||||
StageOutcome::Error.record(&span);
|
||||
ErrorCategory::Internal.record(&span);
|
||||
}
|
||||
}
|
||||
result
|
||||
}
|
||||
@@ -0,0 +1,149 @@
|
||||
use std::{
|
||||
io,
|
||||
sync::{Arc, Mutex},
|
||||
};
|
||||
|
||||
use axum::body::to_bytes;
|
||||
use crank_core::InvocationStatus;
|
||||
use crank_observability::{
|
||||
ObservabilityConfig, OperationalIncident, RedactionLimits, ServiceIdentity,
|
||||
operational_incident_total,
|
||||
};
|
||||
use crank_registry::{
|
||||
InvocationHistoryLoss, InvocationHistoryLossCategory, InvocationHistoryWriteOutcome,
|
||||
};
|
||||
use serde_json::{Value, json};
|
||||
use tracing_subscriber::fmt::MakeWriter;
|
||||
|
||||
use super::{
|
||||
ResponseMode, metrics::normalized_mcp_method, observe_invocation_history_outcome,
|
||||
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": "Проверьте параметры вызова инструмента."
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn emits_bounded_history_loss_incident() {
|
||||
let writer = SharedLogWriter::default();
|
||||
let subscriber = crank_observability::build_subscriber(
|
||||
ObservabilityConfig::new(
|
||||
ServiceIdentity::try_new("mcp-server", "test", "test").unwrap(),
|
||||
"info",
|
||||
RedactionLimits::default(),
|
||||
),
|
||||
writer.clone(),
|
||||
)
|
||||
.unwrap();
|
||||
let before = operational_incident_total(OperationalIncident::InvocationHistoryLost);
|
||||
let dispatch = tracing::Dispatch::new(subscriber);
|
||||
let _guard = tracing::dispatcher::set_default(&dispatch);
|
||||
|
||||
observe_invocation_history_outcome(
|
||||
InvocationHistoryWriteOutcome::Lost(InvocationHistoryLoss {
|
||||
category: InvocationHistoryLossCategory::Unavailable,
|
||||
}),
|
||||
Some("req_mcp_dc08"),
|
||||
InvocationStatus::Ok,
|
||||
"agent_tool_call",
|
||||
);
|
||||
|
||||
let output = writer.output();
|
||||
assert!(!output.contains("dc08-canary-secret"));
|
||||
let event: Value = output
|
||||
.lines()
|
||||
.map(|line| serde_json::from_str(line).unwrap())
|
||||
.find(|event: &Value| event["event"] == "mcp.invocation_history.lost")
|
||||
.unwrap();
|
||||
assert_eq!(event["request_id"], "req_mcp_dc08");
|
||||
assert_eq!(event["fields"]["source"], "agent_tool_call");
|
||||
assert_eq!(event["fields"]["invocation_status"], "ok");
|
||||
assert_eq!(event["fields"]["error_category"], "unavailable");
|
||||
assert!(operational_incident_total(OperationalIncident::InvocationHistoryLost) > before);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mcp_metric_method_is_always_from_a_closed_set() {
|
||||
assert_eq!(
|
||||
normalized_mcp_method(&json!({"jsonrpc": "2.0", "id": 1, "method": "tools/call"})),
|
||||
"tools_call"
|
||||
);
|
||||
assert_eq!(
|
||||
normalized_mcp_method(
|
||||
&json!({"jsonrpc": "2.0", "id": 2, "method": "customer-controlled-method"})
|
||||
),
|
||||
"unsupported"
|
||||
);
|
||||
assert_eq!(
|
||||
normalized_mcp_method(
|
||||
&json!({"jsonrpc": "2.0", "method": "customer-controlled-notification"})
|
||||
),
|
||||
"notification"
|
||||
);
|
||||
}
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
struct SharedLogWriter {
|
||||
buffer: Arc<Mutex<Vec<u8>>>,
|
||||
}
|
||||
|
||||
impl SharedLogWriter {
|
||||
fn output(&self) -> String {
|
||||
String::from_utf8(self.buffer.lock().unwrap().clone()).unwrap()
|
||||
}
|
||||
}
|
||||
|
||||
impl<'a> MakeWriter<'a> for SharedLogWriter {
|
||||
type Writer = SharedLogGuard;
|
||||
|
||||
fn make_writer(&'a self) -> Self::Writer {
|
||||
SharedLogGuard {
|
||||
buffer: Arc::clone(&self.buffer),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct SharedLogGuard {
|
||||
buffer: Arc<Mutex<Vec<u8>>>,
|
||||
}
|
||||
|
||||
impl io::Write for SharedLogGuard {
|
||||
fn write(&mut self, bytes: &[u8]) -> io::Result<usize> {
|
||||
self.buffer.lock().unwrap().extend_from_slice(bytes);
|
||||
Ok(bytes.len())
|
||||
}
|
||||
|
||||
fn flush(&mut self) -> io::Result<()> {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
@@ -5,11 +5,13 @@ use axum::{
|
||||
response::{IntoResponse, Response},
|
||||
};
|
||||
use crank_core::{ApprovalRequestStatus, InvocationLevel, InvocationSource, InvocationStatus};
|
||||
use crank_observability::RequestId;
|
||||
use crank_registry::{ApprovalRequestRecord, FinishApprovalRequest};
|
||||
use crank_runtime::{RuntimeExecutionRequest, RuntimeRequestContext};
|
||||
use crank_trace::{DbOperation, ErrorCategory, Stage, StageOutcome, observe_db_query};
|
||||
use serde_json::json;
|
||||
use time::OffsetDateTime;
|
||||
use tracing::warn;
|
||||
use tracing::{Instrument, warn};
|
||||
|
||||
use crate::{
|
||||
app::{
|
||||
@@ -35,32 +37,81 @@ pub(super) fn spawn_approval_recovery(state: Arc<AppState>) {
|
||||
}
|
||||
|
||||
async fn recover_approved_requests(state: &Arc<AppState>) {
|
||||
fail_interrupted_requests(state).await;
|
||||
|
||||
for _ in 0..32 {
|
||||
let now = OffsetDateTime::now_utc();
|
||||
let approval = match state
|
||||
.registry
|
||||
.claim_next_recoverable_approval_request(
|
||||
now,
|
||||
now - RECOVERY_GRACE,
|
||||
now - EXECUTION_LEASE,
|
||||
)
|
||||
.await
|
||||
let approval = match observe_db_query(
|
||||
DbOperation::ApprovalWrite,
|
||||
state
|
||||
.registry
|
||||
.claim_next_recoverable_approval_request(now, now - RECOVERY_GRACE),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Some(approval)) => approval,
|
||||
Ok(None) => break,
|
||||
Err(error) => {
|
||||
warn!(error = %error, "approval recovery query failed");
|
||||
Err(_) => {
|
||||
warn!(
|
||||
name: "mcp.approval_recovery.query_failed",
|
||||
error_category = "registry",
|
||||
"approval recovery query failed"
|
||||
);
|
||||
break;
|
||||
}
|
||||
};
|
||||
let Some(path) = approval_agent_path(state, &approval).await else {
|
||||
continue;
|
||||
};
|
||||
if execute_approved_request(state, &path, approval)
|
||||
.await
|
||||
.is_err()
|
||||
let recovery_span = Stage::ApprovalRecovery.span();
|
||||
let result = execute_approved_request(state, &path, approval, None)
|
||||
.instrument(recovery_span.clone())
|
||||
.await;
|
||||
match &result {
|
||||
Ok(_) => StageOutcome::Success.record(&recovery_span),
|
||||
Err(_) => {
|
||||
StageOutcome::Error.record(&recovery_span);
|
||||
ErrorCategory::Approval.record(&recovery_span);
|
||||
}
|
||||
}
|
||||
drop(recovery_span);
|
||||
if result.is_err() {
|
||||
warn!(
|
||||
name: "mcp.approval_recovery.execution_failed",
|
||||
error_category = "runtime",
|
||||
"recovered approval execution did not finish"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn fail_interrupted_requests(state: &Arc<AppState>) {
|
||||
for _ in 0..32 {
|
||||
let stale_before = OffsetDateTime::now_utc() - EXECUTION_LEASE;
|
||||
match observe_db_query(
|
||||
DbOperation::ApprovalWrite,
|
||||
state
|
||||
.registry
|
||||
.fail_next_interrupted_approval_request(stale_before),
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!("recovered approval execution did not finish");
|
||||
Ok(Some(approval)) => {
|
||||
warn!(
|
||||
name: "mcp.approval_recovery.interrupted",
|
||||
approval_id = approval.approval.id.as_str(),
|
||||
"interrupted approval execution was not retried because its outcome is unknown"
|
||||
);
|
||||
}
|
||||
Ok(None) => break,
|
||||
Err(_) => {
|
||||
warn!(
|
||||
name: "mcp.approval_recovery.interrupted_query_failed",
|
||||
error_category = "registry",
|
||||
"interrupted approval recovery query failed"
|
||||
);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -76,8 +127,12 @@ async fn approval_agent_path(
|
||||
{
|
||||
Ok(Some(workspace)) => workspace,
|
||||
Ok(None) => return None,
|
||||
Err(error) => {
|
||||
warn!(error = %error, "approval workspace lookup failed");
|
||||
Err(_) => {
|
||||
warn!(
|
||||
name: "mcp.approval_recovery.workspace_lookup_failed",
|
||||
error_category = "registry",
|
||||
"approval workspace lookup failed"
|
||||
);
|
||||
return None;
|
||||
}
|
||||
};
|
||||
@@ -88,8 +143,12 @@ async fn approval_agent_path(
|
||||
{
|
||||
Ok(Some(agent)) => agent,
|
||||
Ok(None) => return None,
|
||||
Err(error) => {
|
||||
warn!(error = %error, "approval agent lookup failed");
|
||||
Err(_) => {
|
||||
warn!(
|
||||
name: "mcp.approval_recovery.agent_lookup_failed",
|
||||
error_category = "registry",
|
||||
"approval agent lookup failed"
|
||||
);
|
||||
return None;
|
||||
}
|
||||
};
|
||||
@@ -103,7 +162,9 @@ pub(super) async fn execute_approved_request(
|
||||
state: &Arc<AppState>,
|
||||
path: &AgentRoutePath,
|
||||
approval: ApprovalRequestRecord,
|
||||
request_id: Option<&str>,
|
||||
) -> Result<ApprovalRequestRecord, Response> {
|
||||
let request_id = RequestId::resolve(request_id).into_string();
|
||||
let tools = state
|
||||
.catalog
|
||||
.list_tools(&path.workspace_slug, &path.agent_slug)
|
||||
@@ -123,18 +184,17 @@ pub(super) async fn execute_approved_request(
|
||||
&approval.approval.request_payload,
|
||||
);
|
||||
let started_at = Instant::now();
|
||||
let runtime_request_context =
|
||||
RuntimeRequestContext::from_request_id(approval.approval.id.as_str().to_owned())
|
||||
.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,
|
||||
)
|
||||
.with_approval_granted();
|
||||
let runtime_request_context = RuntimeRequestContext::from_request_id(request_id.clone())
|
||||
.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,
|
||||
)
|
||||
.with_approval_granted();
|
||||
let resolved_auth =
|
||||
resolve_operation_auth(state, &tool.workspace_id, &operation.execution_config).await;
|
||||
let result = match resolved_auth {
|
||||
@@ -176,11 +236,11 @@ pub(super) async fn execute_approved_request(
|
||||
),
|
||||
};
|
||||
|
||||
if let Err(error) = persist_invocation(
|
||||
persist_invocation(
|
||||
state,
|
||||
&tool,
|
||||
InvocationRecord {
|
||||
request_id: Some(approval.approval.id.as_str()),
|
||||
request_id: Some(&request_id),
|
||||
tool_name: &tool.tool_name,
|
||||
status: invocation_status,
|
||||
level: invocation_level,
|
||||
@@ -192,46 +252,49 @@ pub(super) async fn execute_approved_request(
|
||||
response_preview: response_payload.clone(),
|
||||
},
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(error = %error, "approved invocation log write failed");
|
||||
}
|
||||
.await;
|
||||
|
||||
state
|
||||
.registry
|
||||
.finish_approval_request(FinishApprovalRequest {
|
||||
workspace_id: &approval.approval.workspace_id,
|
||||
agent_id: &approval.approval.agent_id,
|
||||
approval_id: &approval.approval.id,
|
||||
status,
|
||||
response_payload: Some(response_payload),
|
||||
decision_note: None,
|
||||
})
|
||||
.await
|
||||
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR.into_response())?
|
||||
.ok_or_else(|| StatusCode::CONFLICT.into_response())
|
||||
observe_db_query(
|
||||
DbOperation::ApprovalWrite,
|
||||
state
|
||||
.registry
|
||||
.finish_approval_request(FinishApprovalRequest {
|
||||
workspace_id: &approval.approval.workspace_id,
|
||||
agent_id: &approval.approval.agent_id,
|
||||
approval_id: &approval.approval.id,
|
||||
status,
|
||||
response_payload: Some(response_payload),
|
||||
decision_note: None,
|
||||
}),
|
||||
)
|
||||
.await
|
||||
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR.into_response())?
|
||||
.ok_or_else(|| StatusCode::CONFLICT.into_response())
|
||||
}
|
||||
|
||||
async fn finish_unavailable_approval(
|
||||
state: &Arc<AppState>,
|
||||
approval: &ApprovalRequestRecord,
|
||||
) -> Result<ApprovalRequestRecord, Response> {
|
||||
state
|
||||
.registry
|
||||
.finish_approval_request(FinishApprovalRequest {
|
||||
workspace_id: &approval.approval.workspace_id,
|
||||
agent_id: &approval.approval.agent_id,
|
||||
approval_id: &approval.approval.id,
|
||||
status: ApprovalRequestStatus::Failed,
|
||||
response_payload: Some(json!({
|
||||
"error": {
|
||||
"code": "approved_operation_unavailable",
|
||||
"message": "the approved operation version is no longer published"
|
||||
}
|
||||
})),
|
||||
decision_note: None,
|
||||
})
|
||||
.await
|
||||
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR.into_response())?
|
||||
.ok_or_else(|| StatusCode::CONFLICT.into_response())
|
||||
observe_db_query(
|
||||
DbOperation::ApprovalWrite,
|
||||
state
|
||||
.registry
|
||||
.finish_approval_request(FinishApprovalRequest {
|
||||
workspace_id: &approval.approval.workspace_id,
|
||||
agent_id: &approval.approval.agent_id,
|
||||
approval_id: &approval.approval.id,
|
||||
status: ApprovalRequestStatus::Failed,
|
||||
response_payload: Some(json!({
|
||||
"error": {
|
||||
"code": "approved_operation_unavailable",
|
||||
"message": "the approved operation version is no longer published"
|
||||
}
|
||||
})),
|
||||
decision_note: None,
|
||||
}),
|
||||
)
|
||||
.await
|
||||
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR.into_response())?
|
||||
.ok_or_else(|| StatusCode::CONFLICT.into_response())
|
||||
}
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
use crank_core::{ApprovalRequest, OperationApprovalPolicy};
|
||||
use crank_registry::PublishedAgentTool;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
pub(super) fn approval_required_response(
|
||||
tool: &PublishedAgentTool,
|
||||
approval: &ApprovalRequest,
|
||||
policy: &OperationApprovalPolicy,
|
||||
) -> Value {
|
||||
let approval_url = format!(
|
||||
"/v1/{}/{}/approvals/{}",
|
||||
tool.workspace_slug,
|
||||
tool.agent_slug,
|
||||
approval.id.as_str()
|
||||
);
|
||||
json!({
|
||||
"status": "approval_required",
|
||||
"approval_id": approval.id.as_str(),
|
||||
"approval_url": approval_url,
|
||||
"approve": {
|
||||
"method": "POST",
|
||||
"url": format!("{approval_url}/approve"),
|
||||
"body": { "approve": "yes" }
|
||||
},
|
||||
"deny": {
|
||||
"method": "POST",
|
||||
"url": format!("{approval_url}/deny"),
|
||||
"body": { "approve": "no" }
|
||||
},
|
||||
"expires_at": approval.expires_at,
|
||||
"risk_level": approval.risk_level,
|
||||
"payload_preview": if policy.show_payload_preview {
|
||||
approval.request_payload.clone()
|
||||
} else {
|
||||
Value::Null
|
||||
},
|
||||
})
|
||||
}
|
||||
@@ -1,24 +1,27 @@
|
||||
use std::{
|
||||
collections::HashMap,
|
||||
sync::Arc,
|
||||
sync::{Arc, Weak},
|
||||
time::{Duration, Instant, SystemTime, UNIX_EPOCH},
|
||||
};
|
||||
|
||||
use crank_core::{CacheScope, CoordinationStateStore, CoordinationStateValue};
|
||||
use crank_registry::{PostgresRegistry, PublishedAgentCatalog, PublishedAgentTool, RegistryError};
|
||||
use crank_trace::{DbOperation, ErrorCategory, Stage, StageOutcome};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tokio::sync::{Mutex, RwLock};
|
||||
use tracing::{info, warn};
|
||||
use tracing::{Instrument, info, warn};
|
||||
|
||||
use crate::manifest::analyze_published_tool_catalog;
|
||||
|
||||
const MAX_LOCAL_CATALOGS: usize = 1_024;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct PublishedToolCatalog {
|
||||
registry: PostgresRegistry,
|
||||
refresh_interval: Duration,
|
||||
coordination_store: Arc<dyn CoordinationStateStore>,
|
||||
cached: Arc<RwLock<HashMap<CatalogKey, CachedCatalog>>>,
|
||||
refresh_locks: Arc<Mutex<HashMap<CatalogKey, Arc<Mutex<()>>>>>,
|
||||
refresh_locks: Arc<Mutex<HashMap<CatalogKey, Weak<Mutex<()>>>>>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
|
||||
@@ -30,6 +33,14 @@ struct CatalogKey {
|
||||
struct CachedCatalog {
|
||||
loaded_at: Option<Instant>,
|
||||
catalog: PublishedAgentCatalog,
|
||||
metrics: CatalogMetrics,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default)]
|
||||
struct CatalogMetrics {
|
||||
tool_count: usize,
|
||||
estimated_context_tokens: usize,
|
||||
warning_count: usize,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
@@ -66,15 +77,29 @@ impl PublishedToolCatalog {
|
||||
workspace_slug: &str,
|
||||
agent_slug: &str,
|
||||
) -> Result<PublishedAgentCatalog, RegistryError> {
|
||||
self.refresh_if_stale(workspace_slug, agent_slug).await?;
|
||||
let guard = self.cached.read().await;
|
||||
guard
|
||||
.get(&CatalogKey::new(workspace_slug, agent_slug))
|
||||
.map(|entry| entry.catalog.clone())
|
||||
.ok_or_else(|| RegistryError::PublishedAgentNotFound {
|
||||
workspace_slug: workspace_slug.to_owned(),
|
||||
agent_slug: agent_slug.to_owned(),
|
||||
})
|
||||
let span = Stage::McpCatalogLoad.span();
|
||||
let result = async {
|
||||
self.refresh_if_stale(workspace_slug, agent_slug).await?;
|
||||
let guard = self.cached.read().await;
|
||||
guard
|
||||
.get(&CatalogKey::new(workspace_slug, agent_slug))
|
||||
.map(|entry| entry.catalog.clone())
|
||||
.ok_or_else(|| RegistryError::PublishedAgentNotFound {
|
||||
workspace_slug: workspace_slug.to_owned(),
|
||||
agent_slug: agent_slug.to_owned(),
|
||||
})
|
||||
}
|
||||
.instrument(span.clone())
|
||||
.await;
|
||||
match &result {
|
||||
Ok(_) => StageOutcome::Success.record(&span),
|
||||
Err(_) => {
|
||||
StageOutcome::Error.record(&span);
|
||||
ErrorCategory::Catalog.record(&span);
|
||||
}
|
||||
}
|
||||
drop(span);
|
||||
result
|
||||
}
|
||||
|
||||
async fn refresh_if_stale(
|
||||
@@ -98,11 +123,14 @@ impl PublishedToolCatalog {
|
||||
|
||||
let refresh_lock = {
|
||||
let mut locks = self.refresh_locks.lock().await;
|
||||
Arc::clone(
|
||||
locks
|
||||
.entry(key.clone())
|
||||
.or_insert_with(|| Arc::new(Mutex::new(()))),
|
||||
)
|
||||
locks.retain(|_, lock| lock.strong_count() > 0);
|
||||
if let Some(lock) = locks.get(&key).and_then(Weak::upgrade) {
|
||||
lock
|
||||
} else {
|
||||
let lock = Arc::new(Mutex::new(()));
|
||||
locks.insert(key.clone(), Arc::downgrade(&lock));
|
||||
lock
|
||||
}
|
||||
};
|
||||
let _refresh_guard = refresh_lock.lock().await;
|
||||
let still_stale = {
|
||||
@@ -117,50 +145,59 @@ impl PublishedToolCatalog {
|
||||
}
|
||||
|
||||
if let Some((catalog, age)) = self.load_shared_snapshot(workspace_slug, agent_slug).await {
|
||||
log_catalog_analysis(workspace_slug, agent_slug, "shared_cache", &catalog.tools);
|
||||
let mut guard = self.cached.write().await;
|
||||
guard.insert(
|
||||
let metrics =
|
||||
log_catalog_analysis(workspace_slug, agent_slug, "shared_cache", &catalog.tools);
|
||||
self.store_local_catalog(
|
||||
key,
|
||||
CachedCatalog {
|
||||
loaded_at: Instant::now().checked_sub(age),
|
||||
catalog,
|
||||
metrics,
|
||||
},
|
||||
);
|
||||
)
|
||||
.await;
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let catalog = match self
|
||||
let db_span = Stage::DbQuery
|
||||
.db_span(DbOperation::CatalogLoad)
|
||||
.expect("database stage");
|
||||
let catalog_result = self
|
||||
.registry
|
||||
.get_published_agent_catalog_by_slug(workspace_slug, agent_slug)
|
||||
.await
|
||||
{
|
||||
.instrument(db_span.clone())
|
||||
.await;
|
||||
let catalog = match catalog_result {
|
||||
Ok(catalog) => catalog,
|
||||
Err(error) => return Err(error),
|
||||
Err(error) => {
|
||||
StageOutcome::Error.record(&db_span);
|
||||
ErrorCategory::Database.record(&db_span);
|
||||
drop(db_span);
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
log_catalog_analysis(workspace_slug, agent_slug, "postgres", &catalog.tools);
|
||||
StageOutcome::Success.record(&db_span);
|
||||
drop(db_span);
|
||||
let metrics = log_catalog_analysis(workspace_slug, agent_slug, "postgres", &catalog.tools);
|
||||
self.store_shared_snapshot(workspace_slug, agent_slug, &catalog)
|
||||
.await;
|
||||
let mut guard = self.cached.write().await;
|
||||
let previous_count = guard
|
||||
.get(&key)
|
||||
.map(|entry| entry.catalog.tools.len())
|
||||
.unwrap_or_default();
|
||||
|
||||
guard.insert(
|
||||
key,
|
||||
CachedCatalog {
|
||||
loaded_at: Some(Instant::now()),
|
||||
catalog,
|
||||
},
|
||||
);
|
||||
let published_tool_count = catalog.tools.len();
|
||||
let previous_count = self
|
||||
.store_local_catalog(
|
||||
key,
|
||||
CachedCatalog {
|
||||
loaded_at: Some(Instant::now()),
|
||||
catalog,
|
||||
metrics,
|
||||
},
|
||||
)
|
||||
.await;
|
||||
|
||||
info!(
|
||||
name: "mcp.catalog.refreshed",
|
||||
workspace_slug,
|
||||
agent_slug,
|
||||
published_tool_count = guard
|
||||
.get(&CatalogKey::new(workspace_slug, agent_slug))
|
||||
.map(|entry| entry.catalog.tools.len())
|
||||
.unwrap_or_default(),
|
||||
published_tool_count,
|
||||
previous_published_tool_count = previous_count,
|
||||
"published agent catalog refreshed"
|
||||
);
|
||||
@@ -168,6 +205,27 @@ impl PublishedToolCatalog {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn store_local_catalog(&self, key: CatalogKey, entry: CachedCatalog) -> usize {
|
||||
let mut guard = self.cached.write().await;
|
||||
let previous_count = guard
|
||||
.get(&key)
|
||||
.map(|current| current.catalog.tools.len())
|
||||
.unwrap_or_default();
|
||||
|
||||
if guard.len() >= MAX_LOCAL_CATALOGS && !guard.contains_key(&key) {
|
||||
let oldest = guard
|
||||
.iter()
|
||||
.min_by_key(|(_, current)| current.loaded_at)
|
||||
.map(|(candidate, _)| candidate.clone());
|
||||
if let Some(oldest) = oldest {
|
||||
guard.remove(&oldest);
|
||||
}
|
||||
}
|
||||
guard.insert(key, entry);
|
||||
record_catalog_metrics(guard.values().map(|entry| entry.metrics));
|
||||
previous_count
|
||||
}
|
||||
|
||||
async fn load_shared_snapshot(
|
||||
&self,
|
||||
workspace_slug: &str,
|
||||
@@ -226,12 +284,19 @@ fn log_catalog_analysis(
|
||||
agent_slug: &str,
|
||||
source: &str,
|
||||
tools: &[PublishedAgentTool],
|
||||
) {
|
||||
) -> CatalogMetrics {
|
||||
let analysis = match analyze_published_tool_catalog(tools) {
|
||||
Ok(analysis) => analysis,
|
||||
Err(error) => {
|
||||
warn!(workspace_slug, agent_slug, source, %error, "published catalog analysis failed");
|
||||
return;
|
||||
Err(_) => {
|
||||
warn!(
|
||||
name: "mcp.catalog.analysis_failed",
|
||||
workspace_slug,
|
||||
agent_slug,
|
||||
source,
|
||||
error_category = "catalog_validation",
|
||||
"published catalog analysis failed"
|
||||
);
|
||||
return CatalogMetrics::default();
|
||||
}
|
||||
};
|
||||
let warning_count = analysis
|
||||
@@ -242,6 +307,7 @@ fn log_catalog_analysis(
|
||||
.count();
|
||||
|
||||
info!(
|
||||
name: "mcp.catalog.analyzed",
|
||||
workspace_slug,
|
||||
agent_slug,
|
||||
source,
|
||||
@@ -255,6 +321,30 @@ fn log_catalog_analysis(
|
||||
catalog_quality_warning_count = warning_count,
|
||||
"published agent catalog analyzed"
|
||||
);
|
||||
|
||||
CatalogMetrics {
|
||||
tool_count: analysis.budget.tool_count,
|
||||
estimated_context_tokens: analysis.budget.estimated_context_tokens,
|
||||
warning_count,
|
||||
}
|
||||
}
|
||||
|
||||
fn record_catalog_metrics(metrics: impl Iterator<Item = CatalogMetrics>) {
|
||||
let aggregate = metrics.fold(CatalogMetrics::default(), |mut aggregate, current| {
|
||||
aggregate.tool_count = aggregate.tool_count.saturating_add(current.tool_count);
|
||||
aggregate.estimated_context_tokens = aggregate
|
||||
.estimated_context_tokens
|
||||
.saturating_add(current.estimated_context_tokens);
|
||||
aggregate.warning_count = aggregate
|
||||
.warning_count
|
||||
.saturating_add(current.warning_count);
|
||||
aggregate
|
||||
});
|
||||
|
||||
metrics::gauge!("crank_catalog_tools").set(aggregate.tool_count as f64);
|
||||
metrics::gauge!("crank_catalog_estimated_context_tokens")
|
||||
.set(aggregate.estimated_context_tokens as f64);
|
||||
metrics::gauge!("crank_catalog_warnings").set(aggregate.warning_count as f64);
|
||||
}
|
||||
|
||||
fn now_unix_ms() -> u64 {
|
||||
|
||||
@@ -1,14 +1,18 @@
|
||||
mod access;
|
||||
mod app;
|
||||
mod approval_execution;
|
||||
mod approval_response;
|
||||
pub mod auth;
|
||||
pub mod catalog;
|
||||
pub mod jsonrpc;
|
||||
pub mod manifest;
|
||||
mod rate_limit;
|
||||
mod request_context;
|
||||
pub mod session;
|
||||
pub mod tool_error;
|
||||
mod tool_search;
|
||||
mod transport;
|
||||
|
||||
pub use app::{build_app, build_app_with_background_workers};
|
||||
pub use app::{
|
||||
build_app, build_app_with_background_workers, build_app_with_background_workers_and_limits,
|
||||
};
|
||||
|
||||
@@ -4,7 +4,7 @@ use axum::{
|
||||
http::{HeaderMap, HeaderValue, StatusCode, header::RETRY_AFTER},
|
||||
response::{IntoResponse, Response},
|
||||
};
|
||||
use crank_runtime::RateLimitRejection;
|
||||
use crank_runtime::RateLimitCheckError;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use crate::{
|
||||
@@ -14,19 +14,11 @@ use crate::{
|
||||
transport::{ResponseMode, session_id_from_headers, transport_response},
|
||||
};
|
||||
|
||||
pub(super) async fn enforce_post_rate_limit(
|
||||
state: &Arc<AppState>,
|
||||
path: &AgentRoutePath,
|
||||
headers: &HeaderMap,
|
||||
) -> Result<(), RateLimitRejection> {
|
||||
enforce_transport_rate_limit(state, path, headers).await
|
||||
}
|
||||
|
||||
pub(super) async fn enforce_transport_rate_limit(
|
||||
state: &Arc<AppState>,
|
||||
path: &AgentRoutePath,
|
||||
headers: &HeaderMap,
|
||||
) -> Result<(), RateLimitRejection> {
|
||||
) -> Result<(), RateLimitCheckError> {
|
||||
let key = rate_limit_key(path, headers);
|
||||
state.api_rate_limiter.check(&key).await
|
||||
}
|
||||
@@ -35,8 +27,25 @@ pub(super) fn rate_limited_jsonrpc_response(
|
||||
message: &Value,
|
||||
response_mode: ResponseMode,
|
||||
protocol_version: &str,
|
||||
rejection: RateLimitRejection,
|
||||
error: RateLimitCheckError,
|
||||
) -> Response {
|
||||
let RateLimitCheckError::Rejected(rejection) = error else {
|
||||
return transport_response(
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
json!({
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id(message),
|
||||
"error": {
|
||||
"code": -32603,
|
||||
"message": "rate limit service unavailable",
|
||||
"data": { "code": "rate_limit_unavailable" }
|
||||
}
|
||||
}),
|
||||
response_mode,
|
||||
None,
|
||||
Some(protocol_version),
|
||||
);
|
||||
};
|
||||
let payload = json!({
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id(message),
|
||||
@@ -61,10 +70,15 @@ pub(super) fn rate_limited_jsonrpc_response(
|
||||
response
|
||||
}
|
||||
|
||||
pub(super) fn rate_limited_status_response(rejection: RateLimitRejection) -> Response {
|
||||
let mut response = StatusCode::TOO_MANY_REQUESTS.into_response();
|
||||
attach_retry_after_header(&mut response, rejection.retry_after_ms);
|
||||
response
|
||||
pub(super) fn rate_limited_status_response(error: RateLimitCheckError) -> Response {
|
||||
match error {
|
||||
RateLimitCheckError::Rejected(rejection) => {
|
||||
let mut response = StatusCode::TOO_MANY_REQUESTS.into_response();
|
||||
attach_retry_after_header(&mut response, rejection.retry_after_ms);
|
||||
response
|
||||
}
|
||||
RateLimitCheckError::StoreUnavailable => StatusCode::SERVICE_UNAVAILABLE.into_response(),
|
||||
}
|
||||
}
|
||||
|
||||
fn rate_limit_key(path: &AgentRoutePath, headers: &HeaderMap) -> String {
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
use axum::{extract::Request, http::HeaderValue, middleware::Next, response::Response};
|
||||
use crank_observability::{RequestId, set_remote_trace_parent, with_request_correlation};
|
||||
use tracing::{Instrument, info_span};
|
||||
|
||||
use crate::transport::HEADER_X_REQUEST_ID;
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(super) struct RequestContext {
|
||||
pub(super) request_id: String,
|
||||
}
|
||||
|
||||
pub(super) async fn apply_request_context(mut request: Request, next: Next) -> Response {
|
||||
let request_id = RequestId::resolve(
|
||||
request
|
||||
.headers()
|
||||
.get(&HEADER_X_REQUEST_ID)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
)
|
||||
.into_string();
|
||||
let context = RequestContext {
|
||||
request_id: request_id.clone(),
|
||||
};
|
||||
let span = info_span!(
|
||||
target: "crank::trace",
|
||||
"mcp.request",
|
||||
request_id = %request_id,
|
||||
);
|
||||
set_remote_trace_parent(&span, request.headers());
|
||||
request.extensions_mut().insert(context);
|
||||
|
||||
with_request_correlation(request_id.clone(), async move {
|
||||
let mut response = next.run(request).instrument(span).await;
|
||||
if let Ok(value) = HeaderValue::from_str(&request_id) {
|
||||
response.headers_mut().insert(HEADER_X_REQUEST_ID, value);
|
||||
}
|
||||
response
|
||||
})
|
||||
.await
|
||||
}
|
||||
@@ -52,6 +52,8 @@ pub trait TransportSessionStore: Send + Sync {
|
||||
) -> Result<bool, SessionStoreError>;
|
||||
|
||||
async fn delete(&self, session_id: &str) -> Result<bool, SessionStoreError>;
|
||||
|
||||
async fn cleanup_expired(&self, now: OffsetDateTime) -> Result<u64, SessionStoreError>;
|
||||
}
|
||||
|
||||
pub type SharedSessionStore = Arc<dyn TransportSessionStore>;
|
||||
@@ -62,6 +64,11 @@ pub struct PostgresTransportSessionStore {
|
||||
}
|
||||
|
||||
impl PostgresTransportSessionStore {
|
||||
pub async fn from_pool(pool: PgPool) -> Result<Self, SessionStoreError> {
|
||||
apply_postgres_migrations(&pool).await?;
|
||||
Ok(Self { pool })
|
||||
}
|
||||
|
||||
pub async fn connect_with_options_and_pool_config(
|
||||
connect_options: PgConnectOptions,
|
||||
pool_config: PostgresPoolConfig,
|
||||
@@ -84,9 +91,7 @@ impl PostgresTransportSessionStore {
|
||||
details: error.to_string(),
|
||||
})?;
|
||||
|
||||
apply_postgres_migrations(&pool).await?;
|
||||
|
||||
Ok(Self { pool })
|
||||
Self::from_pool(pool).await
|
||||
}
|
||||
}
|
||||
|
||||
@@ -164,6 +169,13 @@ impl TransportSessionStore for InMemorySessionStore {
|
||||
let mut guard = self.inner.write().await;
|
||||
Ok(guard.remove(session_id).is_some())
|
||||
}
|
||||
|
||||
async fn cleanup_expired(&self, now: OffsetDateTime) -> Result<u64, SessionStoreError> {
|
||||
let mut guard = self.inner.write().await;
|
||||
let before = guard.len();
|
||||
guard.retain(|_, session| !is_expired(session, now));
|
||||
Ok(u64::try_from(before.saturating_sub(guard.len())).unwrap_or(u64::MAX))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -292,9 +304,68 @@ impl TransportSessionStore for PostgresTransportSessionStore {
|
||||
|
||||
Ok(result.rows_affected() > 0)
|
||||
}
|
||||
|
||||
async fn cleanup_expired(&self, now: OffsetDateTime) -> Result<u64, SessionStoreError> {
|
||||
let result = query(
|
||||
"delete from mcp_transport_sessions
|
||||
where expires_at is not null and expires_at <= $1::timestamptz",
|
||||
)
|
||||
.bind(now)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_err(|error| SessionStoreError {
|
||||
details: error.to_string(),
|
||||
})?;
|
||||
|
||||
Ok(result.rows_affected())
|
||||
}
|
||||
}
|
||||
|
||||
async fn apply_postgres_migrations(pool: &PgPool) -> Result<(), SessionStoreError> {
|
||||
let mut transaction = pool.begin().await.map_err(|error| SessionStoreError {
|
||||
details: error.to_string(),
|
||||
})?;
|
||||
query("select pg_advisory_xact_lock($1)")
|
||||
.bind(0x4352_414E_4B4D_4350_i64)
|
||||
.execute(&mut *transaction)
|
||||
.await
|
||||
.map_err(|error| SessionStoreError {
|
||||
details: error.to_string(),
|
||||
})?;
|
||||
query(
|
||||
"create table if not exists __crank_mcp_migrations (
|
||||
version integer primary key,
|
||||
checksum text not null,
|
||||
applied_at timestamptz not null default now()
|
||||
)",
|
||||
)
|
||||
.execute(&mut *transaction)
|
||||
.await
|
||||
.map_err(|error| SessionStoreError {
|
||||
details: error.to_string(),
|
||||
})?;
|
||||
let applied = query("select checksum from __crank_mcp_migrations where version = 1")
|
||||
.fetch_optional(&mut *transaction)
|
||||
.await
|
||||
.map_err(|error| SessionStoreError {
|
||||
details: error.to_string(),
|
||||
})?;
|
||||
if let Some(row) = applied {
|
||||
let checksum = row.get::<String, _>("checksum");
|
||||
if checksum != "mcp-transport-sessions-v1" {
|
||||
return Err(SessionStoreError {
|
||||
details: format!("modified MCP migration version 1: {checksum}"),
|
||||
});
|
||||
}
|
||||
transaction
|
||||
.commit()
|
||||
.await
|
||||
.map_err(|error| SessionStoreError {
|
||||
details: error.to_string(),
|
||||
})?;
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
query(
|
||||
"create table if not exists mcp_transport_sessions (
|
||||
id text primary key,
|
||||
@@ -308,14 +379,14 @@ async fn apply_postgres_migrations(pool: &PgPool) -> Result<(), SessionStoreErro
|
||||
expires_at timestamptz null
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut *transaction)
|
||||
.await
|
||||
.map_err(|error| SessionStoreError {
|
||||
details: error.to_string(),
|
||||
})?;
|
||||
|
||||
query("alter table mcp_transport_sessions add column if not exists supports_elicitation boolean not null default false")
|
||||
.execute(pool)
|
||||
.execute(&mut *transaction)
|
||||
.await
|
||||
.map_err(|error| SessionStoreError {
|
||||
details: error.to_string(),
|
||||
@@ -324,7 +395,7 @@ async fn apply_postgres_migrations(pool: &PgPool) -> Result<(), SessionStoreErro
|
||||
query(
|
||||
"alter table mcp_transport_sessions add column if not exists expires_at timestamptz null",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut *transaction)
|
||||
.await
|
||||
.map_err(|error| SessionStoreError {
|
||||
details: error.to_string(),
|
||||
@@ -334,12 +405,37 @@ async fn apply_postgres_migrations(pool: &PgPool) -> Result<(), SessionStoreErro
|
||||
"create index if not exists mcp_transport_sessions_workspace_agent_idx
|
||||
on mcp_transport_sessions(workspace_slug, agent_slug, updated_at desc)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut *transaction)
|
||||
.await
|
||||
.map_err(|error| SessionStoreError {
|
||||
details: error.to_string(),
|
||||
})?;
|
||||
|
||||
query(
|
||||
"create index if not exists mcp_transport_sessions_expires_at_idx
|
||||
on mcp_transport_sessions(expires_at)
|
||||
where expires_at is not null",
|
||||
)
|
||||
.execute(&mut *transaction)
|
||||
.await
|
||||
.map_err(|error| SessionStoreError {
|
||||
details: error.to_string(),
|
||||
})?;
|
||||
|
||||
query("insert into __crank_mcp_migrations (version, checksum) values (1, $1)")
|
||||
.bind("mcp-transport-sessions-v1")
|
||||
.execute(&mut *transaction)
|
||||
.await
|
||||
.map_err(|error| SessionStoreError {
|
||||
details: error.to_string(),
|
||||
})?;
|
||||
transaction
|
||||
.commit()
|
||||
.await
|
||||
.map_err(|error| SessionStoreError {
|
||||
details: error.to_string(),
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
@@ -86,6 +86,10 @@ pub fn runtime_error_code(error: &RuntimeError) -> &'static str {
|
||||
RuntimeError::ConfirmationRequired { .. } => "confirmation_required",
|
||||
RuntimeError::InvalidConfirmationToken { .. } => "invalid_confirmation_token",
|
||||
RuntimeError::ConfirmationStoreUnavailable { .. } => "confirmation_unavailable",
|
||||
RuntimeError::IdempotencyStoreUnavailable { .. } => "idempotency_unavailable",
|
||||
RuntimeError::IdempotencyInProgress { .. } => "idempotency_in_progress",
|
||||
RuntimeError::IdempotencyConflict { .. } => "idempotency_conflict",
|
||||
RuntimeError::IdempotencyOutcomeUnknown { .. } => "idempotency_outcome_unknown",
|
||||
RuntimeError::MissingAuthProfile { .. } => "auth_profile_not_found",
|
||||
RuntimeError::MissingSecret { .. } | RuntimeError::MissingSecretVersion { .. } => {
|
||||
"secret_not_found"
|
||||
@@ -146,6 +150,19 @@ fn safe_runtime_error_message(error: &RuntimeError) -> String {
|
||||
RuntimeError::ConfirmationStoreUnavailable { .. } => {
|
||||
"Хранилище подтверждений временно недоступно.".to_owned()
|
||||
}
|
||||
RuntimeError::IdempotencyStoreUnavailable { .. } => {
|
||||
"Хранилище идемпотентности временно недоступно.".to_owned()
|
||||
}
|
||||
RuntimeError::IdempotencyInProgress { .. } => {
|
||||
"Операция с этим ключом идемпотентности уже выполняется.".to_owned()
|
||||
}
|
||||
RuntimeError::IdempotencyConflict { .. } => {
|
||||
"Ключ идемпотентности уже использован с другими параметрами.".to_owned()
|
||||
}
|
||||
RuntimeError::IdempotencyOutcomeUnknown { .. } => {
|
||||
"Результат предыдущего выполнения неизвестен; автоматический повтор заблокирован."
|
||||
.to_owned()
|
||||
}
|
||||
RuntimeError::MissingAuthProfile { .. } => "Профиль авторизации не найден.".to_owned(),
|
||||
RuntimeError::MissingSecret { .. } | RuntimeError::MissingSecretVersion { .. } => {
|
||||
"Секрет авторизации не найден.".to_owned()
|
||||
@@ -169,6 +186,8 @@ fn is_recoverable_runtime_error(error: &RuntimeError) -> bool {
|
||||
| RuntimeError::ConcurrencyLimitExceeded { .. }
|
||||
| RuntimeError::SecretCrypto { .. }
|
||||
| RuntimeError::ConfirmationRequired { .. }
|
||||
| RuntimeError::IdempotencyStoreUnavailable { .. }
|
||||
| RuntimeError::IdempotencyInProgress { .. }
|
||||
)
|
||||
}
|
||||
|
||||
@@ -198,6 +217,14 @@ fn suggested_action(error: &RuntimeError) -> Option<&'static str> {
|
||||
Some("Запросите новый токен подтверждения.")
|
||||
}
|
||||
RuntimeError::ConfirmationStoreUnavailable { .. } => Some("Повторите запрос позже."),
|
||||
RuntimeError::IdempotencyStoreUnavailable { .. }
|
||||
| RuntimeError::IdempotencyInProgress { .. } => Some("Повторите запрос позже."),
|
||||
RuntimeError::IdempotencyConflict { .. } => {
|
||||
Some("Используйте новый ключ идемпотентности для изменённого запроса.")
|
||||
}
|
||||
RuntimeError::IdempotencyOutcomeUnknown { .. } => {
|
||||
Some("Проверьте результат во внешней системе перед ручным повтором.")
|
||||
}
|
||||
RuntimeError::MissingAuthProfile { .. }
|
||||
| RuntimeError::MissingSecret { .. }
|
||||
| RuntimeError::MissingSecretVersion { .. }
|
||||
|
||||
@@ -3,6 +3,7 @@ use std::{collections::BTreeSet, sync::Arc};
|
||||
use axum::{http::StatusCode, response::Response};
|
||||
use crank_core::{ToolAccessMode, search_tool_catalog};
|
||||
use crank_registry::PublishedAgentCatalog;
|
||||
use crank_trace::{ErrorCategory, Stage, StageOutcome};
|
||||
use serde::Deserialize;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
@@ -127,13 +128,25 @@ async fn execute_catalog_tool(
|
||||
mut arguments: Value,
|
||||
transport_request_id: &str,
|
||||
) -> Response {
|
||||
let Some(resolved) = resolve_generated_tool(&catalog.tools, tool_name) else {
|
||||
return tool_not_found_response(
|
||||
message,
|
||||
response_mode,
|
||||
&session.protocol_version,
|
||||
tool_name,
|
||||
);
|
||||
let resolve_span = Stage::McpToolsResolve.span();
|
||||
let resolved = resolve_span.in_scope(|| resolve_generated_tool(&catalog.tools, tool_name));
|
||||
let resolved = match resolved {
|
||||
Some(resolved) => {
|
||||
StageOutcome::Success.record(&resolve_span);
|
||||
drop(resolve_span);
|
||||
resolved
|
||||
}
|
||||
None => {
|
||||
StageOutcome::Error.record(&resolve_span);
|
||||
ErrorCategory::Catalog.record(&resolve_span);
|
||||
drop(resolve_span);
|
||||
return tool_not_found_response(
|
||||
message,
|
||||
response_mode,
|
||||
&session.protocol_version,
|
||||
tool_name,
|
||||
);
|
||||
}
|
||||
};
|
||||
let confirmation_token = take_confirmation_token(&mut arguments);
|
||||
handle_tool_call(
|
||||
|
||||
@@ -22,7 +22,6 @@ use crate::jsonrpc::{
|
||||
pub(super) const HEADER_MCP_SESSION_ID: &str = "MCP-Session-Id";
|
||||
pub(super) const HEADER_MCP_PROTOCOL_VERSION: &str = "MCP-Protocol-Version";
|
||||
pub(super) const HEADER_X_REQUEST_ID: HeaderName = HeaderName::from_static("x-request-id");
|
||||
const MAX_REQUEST_ID_LEN: usize = 128;
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
pub(super) enum ResponseMode {
|
||||
@@ -303,24 +302,6 @@ where
|
||||
response
|
||||
}
|
||||
|
||||
pub(super) fn resolve_request_id(headers: &HeaderMap) -> String {
|
||||
headers
|
||||
.get(&HEADER_X_REQUEST_ID)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(str::trim)
|
||||
.filter(|value| is_valid_request_id(value))
|
||||
.map(ToOwned::to_owned)
|
||||
.unwrap_or_else(|| uuid::Uuid::now_v7().to_string())
|
||||
}
|
||||
|
||||
pub(super) fn is_valid_request_id(value: &str) -> bool {
|
||||
!value.is_empty()
|
||||
&& value.len() <= MAX_REQUEST_ID_LEN
|
||||
&& value
|
||||
.bytes()
|
||||
.all(|byte| matches!(byte, 0x21..=0x7e) && byte != b',' && byte != b';')
|
||||
}
|
||||
|
||||
pub(super) fn extract_origin(url: &str) -> Option<String> {
|
||||
Some(parse_origin(url, false)?.origin().ascii_serialization())
|
||||
}
|
||||
|
||||
@@ -97,3 +97,41 @@ async fn postgres_transport_sessions_evict_expired_rows_on_read() {
|
||||
|
||||
assert_eq!(remaining, 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn postgres_transport_session_cleanup_removes_abandoned_expired_rows() {
|
||||
let database_url = crank_test_support::postgres_schema_url("test_mcp_cleanup").await;
|
||||
let store = PostgresTransportSessionStore::connect_with_options_and_pool_config(
|
||||
database_url.parse::<PgConnectOptions>().unwrap(),
|
||||
PostgresPoolConfig::default(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let now = OffsetDateTime::now_utc();
|
||||
|
||||
store
|
||||
.create(
|
||||
"2025-11-25",
|
||||
"default",
|
||||
"sales",
|
||||
false,
|
||||
now - time::Duration::hours(2),
|
||||
Some(now - time::Duration::hours(1)),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let active = store
|
||||
.create(
|
||||
"2025-11-25",
|
||||
"default",
|
||||
"sales",
|
||||
false,
|
||||
now,
|
||||
Some(now + time::Duration::hours(1)),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(store.cleanup_expired(now).await.unwrap(), 1);
|
||||
assert!(store.get(&active).await.unwrap().is_some());
|
||||
}
|
||||
|
||||
@@ -72,6 +72,31 @@ async fn drops_expired_in_memory_transport_sessions_on_read() {
|
||||
assert!(store.get(&session_id).await.unwrap().is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cleanup_removes_only_expired_sessions() {
|
||||
let store = InMemorySessionStore::default();
|
||||
let now = time::OffsetDateTime::now_utc();
|
||||
let expired = store
|
||||
.create("2025-11-25", "default", "sales", false, now, Some(now))
|
||||
.await
|
||||
.unwrap();
|
||||
let active = store
|
||||
.create(
|
||||
"2025-11-25",
|
||||
"default",
|
||||
"sales",
|
||||
false,
|
||||
now,
|
||||
Some(now + time::Duration::hours(1)),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(store.cleanup_expired(now).await.unwrap(), 1);
|
||||
assert!(store.get(&expired).await.unwrap().is_none());
|
||||
assert!(store.get(&active).await.unwrap().is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn formats_transport_session_store_error() {
|
||||
let error = SessionStoreError {
|
||||
|
||||
@@ -3,6 +3,7 @@ name = "crank-core"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
rust-version.workspace = true
|
||||
publish.workspace = true
|
||||
version.workspace = true
|
||||
|
||||
[dependencies]
|
||||
|
||||
@@ -74,6 +74,12 @@ pub struct RateLimitBucketState {
|
||||
pub last_refill_unix_ms: i64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum RateLimitDecision {
|
||||
Allowed,
|
||||
Rejected { retry_after_ms: u64 },
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum ReplayGuardStatus {
|
||||
@@ -86,6 +92,12 @@ pub struct CoordinationStateValue {
|
||||
pub payload: Value,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum CoordinationStateReservation {
|
||||
Reserved,
|
||||
Existing(CoordinationStateValue),
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait ResponseCacheStore: Send + Sync {
|
||||
async fn get(&self, key: &str) -> Result<Option<CachedResponse>, CacheStoreError>;
|
||||
@@ -108,6 +120,14 @@ pub trait RateLimitStateStore: Send + Sync {
|
||||
ttl: Duration,
|
||||
) -> Result<(), CacheStoreError>;
|
||||
async fn delete_bucket(&self, key: &str) -> Result<(), CacheStoreError>;
|
||||
async fn consume_token(
|
||||
&self,
|
||||
key: &str,
|
||||
burst_tokens_micros: u64,
|
||||
refill_per_second_micros: u64,
|
||||
now_unix_ms: i64,
|
||||
ttl: Duration,
|
||||
) -> Result<RateLimitDecision, CacheStoreError>;
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -135,6 +155,26 @@ pub trait CoordinationStateStore: Send + Sync {
|
||||
ttl: Duration,
|
||||
) -> Result<(), CacheStoreError>;
|
||||
async fn delete_value(&self, scope: CacheScope, key: &str) -> Result<(), CacheStoreError>;
|
||||
async fn take_value(
|
||||
&self,
|
||||
scope: CacheScope,
|
||||
key: &str,
|
||||
) -> Result<Option<CoordinationStateValue>, CacheStoreError>;
|
||||
async fn reserve_value(
|
||||
&self,
|
||||
scope: CacheScope,
|
||||
key: &str,
|
||||
value: CoordinationStateValue,
|
||||
ttl: Duration,
|
||||
) -> Result<CoordinationStateReservation, CacheStoreError>;
|
||||
async fn compare_and_set_value(
|
||||
&self,
|
||||
scope: CacheScope,
|
||||
key: &str,
|
||||
expected: &CoordinationStateValue,
|
||||
value: CoordinationStateValue,
|
||||
ttl: Duration,
|
||||
) -> Result<bool, CacheStoreError>;
|
||||
}
|
||||
|
||||
#[derive(Debug, Error, PartialEq, Eq)]
|
||||
|
||||
@@ -66,8 +66,9 @@ pub mod domain {
|
||||
|
||||
pub mod ports {
|
||||
pub use crate::cache::{
|
||||
CacheStoreError, CoordinationStateStore, CoordinationStateValue, RateLimitStateStore,
|
||||
ReplayGuardStatus, ReplayGuardStore, ResponseCacheStore,
|
||||
CacheStoreError, CoordinationStateReservation, CoordinationStateStore,
|
||||
CoordinationStateValue, RateLimitDecision, RateLimitStateStore, ReplayGuardStatus,
|
||||
ReplayGuardStore, ResponseCacheStore,
|
||||
};
|
||||
pub use crate::ext::access::{
|
||||
OwnerOnlyPolicyEngine, PolicyAction, PolicyDecision, PolicyEngine, PolicyScope,
|
||||
@@ -108,8 +109,9 @@ pub use auth::{
|
||||
};
|
||||
pub use cache::{
|
||||
CacheBackend, CacheScope, CacheStoreError, CachedHeader, CachedResponse,
|
||||
CoordinationStateStore, CoordinationStateValue, ParseCacheBackendError, RateLimitBucketState,
|
||||
RateLimitStateStore, ReplayGuardStatus, ReplayGuardStore, ResponseCacheStore,
|
||||
CoordinationStateReservation, CoordinationStateStore, CoordinationStateValue,
|
||||
ParseCacheBackendError, RateLimitBucketState, RateLimitDecision, RateLimitStateStore,
|
||||
ReplayGuardStatus, ReplayGuardStore, ResponseCacheStore,
|
||||
};
|
||||
pub use edition::{
|
||||
EditionCapabilities, EditionLimits, MachineAccessMode, OperationSecurityLevel, ProductEdition,
|
||||
|
||||
@@ -3,6 +3,7 @@ name = "crank-import"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
rust-version.workspace = true
|
||||
publish.workspace = true
|
||||
version.workspace = true
|
||||
|
||||
[dependencies]
|
||||
|
||||
@@ -3,6 +3,7 @@ name = "crank-mapping"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
rust-version.workspace = true
|
||||
publish.workspace = true
|
||||
version.workspace = true
|
||||
|
||||
[dependencies]
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
[package]
|
||||
name = "crank-observability"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
rust-version.workspace = true
|
||||
publish.workspace = true
|
||||
version.workspace = true
|
||||
|
||||
[dependencies]
|
||||
axum.workspace = true
|
||||
metrics.workspace = true
|
||||
metrics-exporter-prometheus.workspace = true
|
||||
opentelemetry.workspace = true
|
||||
opentelemetry-otlp.workspace = true
|
||||
opentelemetry_sdk.workspace = true
|
||||
percent-encoding.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
sentry.workspace = true
|
||||
sha2.workspace = true
|
||||
subtle.workspace = true
|
||||
thiserror.workspace = true
|
||||
time.workspace = true
|
||||
tokio = { workspace = true, features = ["net"] }
|
||||
tracing.workspace = true
|
||||
tracing-opentelemetry.workspace = true
|
||||
tracing-subscriber.workspace = true
|
||||
url.workspace = true
|
||||
uuid.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
opentelemetry-proto.workspace = true
|
||||
prost.workspace = true
|
||||
sentry = { workspace = true, features = ["test"] }
|
||||
tower.workspace = true
|
||||
@@ -0,0 +1,170 @@
|
||||
use std::env;
|
||||
|
||||
use thiserror::Error;
|
||||
|
||||
use crate::RedactionLimits;
|
||||
|
||||
const DEFAULT_ENVIRONMENT: &str = "development";
|
||||
const MAX_IDENTITY_LABEL_BYTES: usize = 64;
|
||||
|
||||
#[derive(Clone, Debug, Eq, PartialEq)]
|
||||
pub struct ServiceIdentity {
|
||||
service: String,
|
||||
version: String,
|
||||
environment: String,
|
||||
}
|
||||
|
||||
impl ServiceIdentity {
|
||||
pub fn try_new(
|
||||
service: impl Into<String>,
|
||||
version: impl Into<String>,
|
||||
environment: impl Into<String>,
|
||||
) -> Result<Self, ObservabilityConfigError> {
|
||||
let identity = Self {
|
||||
service: service.into(),
|
||||
version: version.into(),
|
||||
environment: environment.into(),
|
||||
};
|
||||
validate_label("service", &identity.service)?;
|
||||
validate_label("version", &identity.version)?;
|
||||
validate_label("environment", &identity.environment)?;
|
||||
Ok(identity)
|
||||
}
|
||||
|
||||
pub fn service(&self) -> &str {
|
||||
&self.service
|
||||
}
|
||||
|
||||
pub fn version(&self) -> &str {
|
||||
&self.version
|
||||
}
|
||||
|
||||
pub fn environment(&self) -> &str {
|
||||
&self.environment
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct ObservabilityConfig {
|
||||
identity: ServiceIdentity,
|
||||
filter: String,
|
||||
redaction_limits: RedactionLimits,
|
||||
}
|
||||
|
||||
impl ObservabilityConfig {
|
||||
pub fn new(
|
||||
identity: ServiceIdentity,
|
||||
filter: impl Into<String>,
|
||||
redaction_limits: RedactionLimits,
|
||||
) -> Self {
|
||||
Self {
|
||||
identity,
|
||||
filter: filter.into(),
|
||||
redaction_limits,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn from_env(
|
||||
service: &'static str,
|
||||
version: &'static str,
|
||||
default_filter: &'static str,
|
||||
) -> Result<Self, ObservabilityConfigError> {
|
||||
let environment = env_value_or_default(
|
||||
"CRANK_ENVIRONMENT",
|
||||
env::var("CRANK_ENVIRONMENT"),
|
||||
DEFAULT_ENVIRONMENT,
|
||||
)?;
|
||||
let filter = env_value_or_default(
|
||||
"CRANK_LOG_LEVEL",
|
||||
env::var("CRANK_LOG_LEVEL"),
|
||||
default_filter,
|
||||
)?;
|
||||
let identity = ServiceIdentity::try_new(service, version, environment)?;
|
||||
|
||||
Ok(Self::new(identity, filter, RedactionLimits::default()))
|
||||
}
|
||||
|
||||
pub(crate) fn into_parts(self) -> (ServiceIdentity, String, RedactionLimits) {
|
||||
(self.identity, self.filter, self.redaction_limits)
|
||||
}
|
||||
|
||||
pub(crate) fn identity(&self) -> &ServiceIdentity {
|
||||
&self.identity
|
||||
}
|
||||
|
||||
pub(crate) fn redaction_limits(&self) -> RedactionLimits {
|
||||
self.redaction_limits
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum ObservabilityConfigError {
|
||||
#[error("invalid observability identity field: {field}")]
|
||||
InvalidIdentity { field: &'static str },
|
||||
#[error("observability environment variable is not valid UTF-8: {field}")]
|
||||
InvalidEnvironmentEncoding { field: &'static str },
|
||||
}
|
||||
|
||||
fn env_value_or_default(
|
||||
field: &'static str,
|
||||
value: Result<String, env::VarError>,
|
||||
default: &'static str,
|
||||
) -> Result<String, ObservabilityConfigError> {
|
||||
match value {
|
||||
Ok(value) => Ok(value),
|
||||
Err(env::VarError::NotPresent) => Ok(default.to_owned()),
|
||||
Err(env::VarError::NotUnicode(_)) => {
|
||||
Err(ObservabilityConfigError::InvalidEnvironmentEncoding { field })
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_label(field: &'static str, value: &str) -> Result<(), ObservabilityConfigError> {
|
||||
let valid = !value.is_empty()
|
||||
&& value.len() <= MAX_IDENTITY_LABEL_BYTES
|
||||
&& value
|
||||
.bytes()
|
||||
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.' | b'+'));
|
||||
|
||||
if valid {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(ObservabilityConfigError::InvalidIdentity { field })
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::ffi::OsString;
|
||||
|
||||
use super::{ObservabilityConfigError, ServiceIdentity, env_value_or_default};
|
||||
|
||||
#[test]
|
||||
fn accepts_release_and_environment_labels() {
|
||||
let identity = ServiceIdentity::try_new("admin-api", "0.3.1+build.7", "production")
|
||||
.expect("identity must be valid");
|
||||
|
||||
assert_eq!(identity.service(), "admin-api");
|
||||
assert_eq!(identity.version(), "0.3.1+build.7");
|
||||
assert_eq!(identity.environment(), "production");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_non_utf8_environment_values() {
|
||||
let error = env_value_or_default(
|
||||
"CRANK_ENVIRONMENT",
|
||||
Err(std::env::VarError::NotUnicode(OsString::from(
|
||||
"invalid-environment",
|
||||
))),
|
||||
"development",
|
||||
)
|
||||
.expect_err("non-UTF-8 values must not be replaced with defaults");
|
||||
|
||||
assert!(matches!(
|
||||
error,
|
||||
ObservabilityConfigError::InvalidEnvironmentEncoding {
|
||||
field: "CRANK_ENVIRONMENT"
|
||||
}
|
||||
));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
use std::fmt;
|
||||
|
||||
use uuid::Uuid;
|
||||
|
||||
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
|
||||
pub struct RequestId(String);
|
||||
|
||||
impl RequestId {
|
||||
pub const MAX_LEN: usize = 128;
|
||||
|
||||
pub fn resolve(candidate: Option<&str>) -> Self {
|
||||
candidate
|
||||
.filter(|value| Self::is_valid(value))
|
||||
.map(|value| Self(value.to_owned()))
|
||||
.unwrap_or_else(|| Self(Uuid::now_v7().to_string()))
|
||||
}
|
||||
|
||||
pub fn is_valid(value: &str) -> bool {
|
||||
!value.is_empty()
|
||||
&& value.len() <= Self::MAX_LEN
|
||||
&& value
|
||||
.bytes()
|
||||
.all(|byte| matches!(byte, 0x21..=0x7e) && byte != b',' && byte != b';')
|
||||
}
|
||||
|
||||
pub fn as_str(&self) -> &str {
|
||||
&self.0
|
||||
}
|
||||
|
||||
pub fn into_string(self) -> String {
|
||||
self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for RequestId {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter.write_str(self.as_str())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,467 @@
|
||||
use std::{borrow::Cow, collections::BTreeMap, env, fmt, future::Future, time::Duration};
|
||||
|
||||
use sentry::{
|
||||
ClientInitGuard, ClientOptions,
|
||||
protocol::{Event, Level},
|
||||
types::Dsn,
|
||||
};
|
||||
use thiserror::Error;
|
||||
|
||||
use crate::{
|
||||
RedactionLimits, ServiceIdentity, propagation::current_trace_id, redaction::truncate_string,
|
||||
};
|
||||
|
||||
const SENTRY_DSN_ENV: &str = "CRANK_SENTRY_DSN";
|
||||
const CRITICAL_ERROR_MESSAGE: &str = "critical error";
|
||||
const SENTRY_SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(2);
|
||||
|
||||
tokio::task_local! {
|
||||
static REQUEST_ID: String;
|
||||
}
|
||||
|
||||
pub struct SentryConfig {
|
||||
dsn: Option<Dsn>,
|
||||
}
|
||||
|
||||
impl SentryConfig {
|
||||
pub fn parse(value: Option<&str>) -> Result<Self, SentryConfigError> {
|
||||
let Some(value) = value.map(str::trim).filter(|value| !value.is_empty()) else {
|
||||
return Ok(Self { dsn: None });
|
||||
};
|
||||
|
||||
let dsn = value
|
||||
.parse::<Dsn>()
|
||||
.map_err(|_| SentryConfigError::InvalidDsn)?;
|
||||
Ok(Self { dsn: Some(dsn) })
|
||||
}
|
||||
|
||||
pub fn from_env() -> Result<Self, SentryConfigError> {
|
||||
match env::var(SENTRY_DSN_ENV) {
|
||||
Ok(value) => Self::parse(Some(&value)),
|
||||
Err(env::VarError::NotPresent) => Self::parse(None),
|
||||
Err(env::VarError::NotUnicode(_)) => Err(SentryConfigError::InvalidEnvironmentEncoding),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn enabled(&self) -> bool {
|
||||
self.dsn.is_some()
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for SentryConfig {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("SentryConfig")
|
||||
.field("enabled", &self.enabled())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum SentryConfigError {
|
||||
#[error("CRANK_SENTRY_DSN is not a valid Sentry DSN")]
|
||||
InvalidDsn,
|
||||
#[error("CRANK_SENTRY_DSN is not valid UTF-8")]
|
||||
InvalidEnvironmentEncoding,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub enum CriticalErrorCategory {
|
||||
Panic,
|
||||
Startup,
|
||||
Internal,
|
||||
DataIntegrity,
|
||||
}
|
||||
|
||||
impl CriticalErrorCategory {
|
||||
pub const fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Panic => "panic",
|
||||
Self::Startup => "startup",
|
||||
Self::Internal => "internal",
|
||||
Self::DataIntegrity => "data_integrity",
|
||||
}
|
||||
}
|
||||
|
||||
fn parse(value: &str) -> Option<Self> {
|
||||
match value {
|
||||
"panic" => Some(Self::Panic),
|
||||
"startup" => Some(Self::Startup),
|
||||
"internal" => Some(Self::Internal),
|
||||
"data_integrity" => Some(Self::DataIntegrity),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn capture_critical_error(category: CriticalErrorCategory) {
|
||||
let mut tags = correlation_tags();
|
||||
tags.insert("category".to_owned(), category.as_str().to_owned());
|
||||
sentry::capture_event(Event {
|
||||
level: Level::Error,
|
||||
message: Some(CRITICAL_ERROR_MESSAGE.to_owned()),
|
||||
fingerprint: Cow::Owned(vec![Cow::Borrowed(category.as_str())]),
|
||||
tags,
|
||||
..Event::default()
|
||||
});
|
||||
}
|
||||
|
||||
pub async fn with_request_correlation<F>(request_id: String, future: F) -> F::Output
|
||||
where
|
||||
F: Future,
|
||||
{
|
||||
REQUEST_ID.scope(request_id, future).await
|
||||
}
|
||||
|
||||
pub(crate) fn init_sentry(
|
||||
identity: &ServiceIdentity,
|
||||
limits: RedactionLimits,
|
||||
config: SentryConfig,
|
||||
) -> Option<ClientInitGuard> {
|
||||
let dsn = config.dsn?;
|
||||
let identity = identity.clone();
|
||||
let options = client_options(identity, limits);
|
||||
Some(sentry::init((dsn, options)))
|
||||
}
|
||||
|
||||
fn client_options(identity: ServiceIdentity, limits: RedactionLimits) -> ClientOptions {
|
||||
let release = identity.version().to_owned();
|
||||
let environment = identity.environment().to_owned();
|
||||
let service = identity.service().to_owned();
|
||||
let sanitizer_identity = identity.clone();
|
||||
|
||||
let mut options = ClientOptions::default();
|
||||
options.release = Some(Cow::Owned(release));
|
||||
options.environment = Some(Cow::Owned(environment));
|
||||
options.server_name = Some(Cow::Owned(service));
|
||||
options.traces_sampling_strategy = sentry::TracesSamplingStrategy::Disabled;
|
||||
options.max_breadcrumbs = 0;
|
||||
options.attach_stacktrace = false;
|
||||
options.send_default_pii = false;
|
||||
options.before_send = Some(std::sync::Arc::new(move |event| {
|
||||
Some(sanitize_event(event, &sanitizer_identity, limits))
|
||||
}));
|
||||
options.shutdown_timeout = SENTRY_SHUTDOWN_TIMEOUT;
|
||||
options.auto_session_tracking = false;
|
||||
options.enable_logs = false;
|
||||
options.enable_metrics = false;
|
||||
options
|
||||
}
|
||||
|
||||
fn sanitize_event(
|
||||
event: Event<'static>,
|
||||
identity: &ServiceIdentity,
|
||||
limits: RedactionLimits,
|
||||
) -> Event<'static> {
|
||||
let category = event
|
||||
.tags
|
||||
.get("category")
|
||||
.and_then(|value| CriticalErrorCategory::parse(value))
|
||||
.unwrap_or_else(|| {
|
||||
if event.exception.is_empty() {
|
||||
CriticalErrorCategory::Internal
|
||||
} else {
|
||||
CriticalErrorCategory::Panic
|
||||
}
|
||||
});
|
||||
let mut tags = correlation_tags()
|
||||
.into_iter()
|
||||
.map(|(key, value)| (key, truncate_string(&value, limits.max_string_bytes)))
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
for key in ["request_id", "trace_id"] {
|
||||
if let Some(value) = event.tags.get(key) {
|
||||
tags.entry(key.to_owned())
|
||||
.or_insert_with(|| truncate_string(value, limits.max_string_bytes));
|
||||
}
|
||||
}
|
||||
tags.insert("service".to_owned(), identity.service().to_owned());
|
||||
tags.insert("category".to_owned(), category.as_str().to_owned());
|
||||
|
||||
enforce_event_budget(
|
||||
Event {
|
||||
event_id: event.event_id,
|
||||
level: Level::Error,
|
||||
fingerprint: Cow::Owned(vec![Cow::Borrowed(category.as_str())]),
|
||||
message: Some(CRITICAL_ERROR_MESSAGE.to_owned()),
|
||||
timestamp: event.timestamp,
|
||||
server_name: Some(Cow::Owned(identity.service().to_owned())),
|
||||
release: Some(Cow::Owned(identity.version().to_owned())),
|
||||
environment: Some(Cow::Owned(identity.environment().to_owned())),
|
||||
tags,
|
||||
..Event::default()
|
||||
},
|
||||
limits.max_event_bytes,
|
||||
)
|
||||
}
|
||||
|
||||
fn correlation_tags() -> BTreeMap<String, String> {
|
||||
let mut tags = BTreeMap::new();
|
||||
if let Ok(request_id) = REQUEST_ID.try_with(Clone::clone) {
|
||||
tags.insert("request_id".to_owned(), request_id);
|
||||
}
|
||||
if let Some(trace_id) = current_trace_id() {
|
||||
tags.insert("trace_id".to_owned(), trace_id);
|
||||
}
|
||||
tags
|
||||
}
|
||||
|
||||
fn enforce_event_budget(mut event: Event<'static>, max_event_bytes: usize) -> Event<'static> {
|
||||
if serialized_event_len(&event) <= max_event_bytes {
|
||||
return event;
|
||||
}
|
||||
|
||||
event.tags.remove("request_id");
|
||||
event.tags.remove("trace_id");
|
||||
event
|
||||
}
|
||||
|
||||
fn serialized_event_len(event: &Event<'_>) -> usize {
|
||||
serde_json::to_vec(event).map_or(usize::MAX, |serialized| serialized.len())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::{
|
||||
collections::BTreeMap,
|
||||
sync::{
|
||||
Arc,
|
||||
atomic::{AtomicUsize, Ordering},
|
||||
},
|
||||
};
|
||||
|
||||
use opentelemetry::trace::TracerProvider as _;
|
||||
use opentelemetry_sdk::trace::SdkTracerProvider;
|
||||
use sentry::{
|
||||
Envelope, Hub,
|
||||
protocol::{Breadcrumb, Context, Exception, Request, User, Value, Values},
|
||||
};
|
||||
|
||||
use super::{
|
||||
CriticalErrorCategory, capture_critical_error, client_options, sanitize_event,
|
||||
with_request_correlation,
|
||||
};
|
||||
use crate::{
|
||||
ObservabilityConfig, RedactionLimits, ServiceIdentity,
|
||||
logging::build_subscriber_with_tracer,
|
||||
};
|
||||
|
||||
fn identity() -> ServiceIdentity {
|
||||
ServiceIdentity::try_new("admin-api", "1.2.3", "test").expect("valid identity")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn client_disables_non_error_telemetry() {
|
||||
let options = client_options(identity(), RedactionLimits::default());
|
||||
|
||||
assert_eq!(options.max_breadcrumbs, 0);
|
||||
assert!(!options.attach_stacktrace);
|
||||
assert!(!options.send_default_pii);
|
||||
assert!(!options.auto_session_tracking);
|
||||
assert!(!options.enable_logs);
|
||||
assert!(!options.enable_metrics);
|
||||
assert!(options.before_send.is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sanitizer_uses_a_strict_allowlist() {
|
||||
let mut tags = BTreeMap::new();
|
||||
tags.insert("category".to_owned(), "data_integrity".to_owned());
|
||||
tags.insert("secret".to_owned(), "must-not-leak".to_owned());
|
||||
let mut contexts = BTreeMap::new();
|
||||
contexts.insert(
|
||||
"secret".to_owned(),
|
||||
Context::Other(BTreeMap::from([(
|
||||
"token".to_owned(),
|
||||
Value::String("must-not-leak".to_owned()),
|
||||
)])),
|
||||
);
|
||||
let event = sentry::protocol::Event {
|
||||
message: Some("password=must-not-leak".to_owned()),
|
||||
request: Some(Request::default()),
|
||||
user: Some(User::default()),
|
||||
breadcrumbs: Values {
|
||||
values: vec![Breadcrumb::default()],
|
||||
},
|
||||
exception: Values {
|
||||
values: vec![Exception {
|
||||
value: Some("must-not-leak".to_owned()),
|
||||
..Exception::default()
|
||||
}],
|
||||
},
|
||||
contexts,
|
||||
extra: BTreeMap::from([(
|
||||
"payload".to_owned(),
|
||||
Value::String("must-not-leak".to_owned()),
|
||||
)]),
|
||||
tags,
|
||||
..sentry::protocol::Event::default()
|
||||
};
|
||||
|
||||
let cleaned = sanitize_event(event, &identity(), RedactionLimits::default());
|
||||
let serialized = serde_json::to_string(&cleaned).expect("serialize event");
|
||||
|
||||
assert_eq!(cleaned.message.as_deref(), Some("critical error"));
|
||||
assert_eq!(
|
||||
cleaned.tags.get("category").map(String::as_str),
|
||||
Some(CriticalErrorCategory::DataIntegrity.as_str())
|
||||
);
|
||||
assert!(cleaned.request.is_none());
|
||||
assert!(cleaned.user.is_none());
|
||||
assert!(cleaned.breadcrumbs.is_empty());
|
||||
assert!(cleaned.exception.is_empty());
|
||||
assert!(cleaned.contexts.is_empty());
|
||||
assert!(cleaned.extra.is_empty());
|
||||
assert!(!serialized.contains("must-not-leak"));
|
||||
assert!(!serialized.contains("password"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sanitizer_honours_total_event_budget() {
|
||||
let limits = RedactionLimits {
|
||||
max_string_bytes: 8 * 1024,
|
||||
max_event_bytes: 512,
|
||||
..RedactionLimits::default()
|
||||
};
|
||||
let event = sentry::protocol::Event {
|
||||
tags: BTreeMap::from([
|
||||
("category".to_owned(), "internal".to_owned()),
|
||||
("request_id".to_owned(), "r".repeat(8 * 1024)),
|
||||
("trace_id".to_owned(), "t".repeat(8 * 1024)),
|
||||
]),
|
||||
..sentry::protocol::Event::default()
|
||||
};
|
||||
|
||||
let cleaned = sanitize_event(event, &identity(), limits);
|
||||
let serialized = serde_json::to_vec(&cleaned).expect("serialize event");
|
||||
|
||||
assert!(serialized.len() <= limits.max_event_bytes);
|
||||
assert_eq!(
|
||||
cleaned.tags.get("category").map(String::as_str),
|
||||
Some("internal")
|
||||
);
|
||||
assert_eq!(
|
||||
cleaned.tags.get("service").map(String::as_str),
|
||||
Some("admin-api")
|
||||
);
|
||||
assert!(!cleaned.tags.contains_key("request_id"));
|
||||
assert!(!cleaned.tags.contains_key("trace_id"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn expected_application_errors_do_not_create_critical_events() {
|
||||
let options =
|
||||
sentry::apply_defaults(client_options(identity(), RedactionLimits::default()));
|
||||
let events = sentry::test::with_captured_events_options(
|
||||
|| tracing::error!("ordinary product error"),
|
||||
options,
|
||||
);
|
||||
|
||||
assert!(events.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn explicit_critical_error_is_correlated_and_sanitized() {
|
||||
let options =
|
||||
sentry::apply_defaults(client_options(identity(), RedactionLimits::default()));
|
||||
let runtime = tokio::runtime::Builder::new_current_thread()
|
||||
.build()
|
||||
.expect("runtime");
|
||||
let provider = SdkTracerProvider::builder().build();
|
||||
let tracer = provider.tracer("critical-error-test");
|
||||
let subscriber = build_subscriber_with_tracer(
|
||||
ObservabilityConfig::new(identity(), "info", RedactionLimits::default()),
|
||||
std::io::sink,
|
||||
Some(tracer),
|
||||
)
|
||||
.expect("subscriber");
|
||||
let dispatch = tracing::Dispatch::new(subscriber);
|
||||
let events = sentry::test::with_captured_events_options(
|
||||
|| {
|
||||
tracing::dispatcher::with_default(&dispatch, || {
|
||||
runtime.block_on(with_request_correlation("request-123".to_owned(), async {
|
||||
let span = tracing::info_span!(target: "crank::trace", "http.request");
|
||||
let _span_guard = span.enter();
|
||||
capture_critical_error(CriticalErrorCategory::DataIntegrity);
|
||||
}));
|
||||
});
|
||||
},
|
||||
options,
|
||||
);
|
||||
|
||||
assert_eq!(events.len(), 1);
|
||||
let event = &events[0];
|
||||
assert_eq!(event.message.as_deref(), Some("critical error"));
|
||||
assert_eq!(
|
||||
event.tags.get("category").map(String::as_str),
|
||||
Some("data_integrity")
|
||||
);
|
||||
assert_eq!(
|
||||
event.tags.get("request_id").map(String::as_str),
|
||||
Some("request-123")
|
||||
);
|
||||
assert_eq!(event.tags.get("trace_id").map(String::len), Some(32));
|
||||
assert_eq!(event.release.as_deref(), Some("1.2.3"));
|
||||
assert_eq!(event.environment.as_deref(), Some("test"));
|
||||
assert_eq!(event.server_name.as_deref(), Some("admin-api"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn panic_creates_exactly_one_sanitized_critical_event() {
|
||||
let options =
|
||||
sentry::apply_defaults(client_options(identity(), RedactionLimits::default()));
|
||||
let events = sentry::test::with_captured_events_options(
|
||||
|| {
|
||||
let result = std::panic::catch_unwind(|| {
|
||||
panic!("password=must-not-leak");
|
||||
});
|
||||
assert!(result.is_err());
|
||||
},
|
||||
options,
|
||||
);
|
||||
|
||||
assert_eq!(events.len(), 1);
|
||||
let event = &events[0];
|
||||
assert_eq!(
|
||||
event.tags.get("category").map(String::as_str),
|
||||
Some("panic")
|
||||
);
|
||||
let serialized = serde_json::to_string(event).expect("serialize event");
|
||||
assert!(!serialized.contains("must-not-leak"));
|
||||
assert!(!serialized.contains("password"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn receiver_failure_does_not_change_product_result_or_recurse() {
|
||||
struct DroppingTransport {
|
||||
attempts: AtomicUsize,
|
||||
}
|
||||
|
||||
impl sentry::Transport for DroppingTransport {
|
||||
fn send_envelope(&self, _envelope: Envelope) {
|
||||
self.attempts.fetch_add(1, Ordering::Relaxed);
|
||||
}
|
||||
}
|
||||
|
||||
let transport = Arc::new(DroppingTransport {
|
||||
attempts: AtomicUsize::new(0),
|
||||
});
|
||||
let mut options =
|
||||
sentry::apply_defaults(client_options(identity(), RedactionLimits::default()));
|
||||
options.dsn = Some(
|
||||
"https://public@example.invalid/1"
|
||||
.parse()
|
||||
.expect("valid test DSN"),
|
||||
);
|
||||
options.transport = Some(Arc::new(transport.clone()));
|
||||
let client = Arc::new(sentry::Client::from(options));
|
||||
let hub = Arc::new(Hub::new(Some(client), Arc::new(Default::default())));
|
||||
|
||||
let product_result = Hub::run(hub, || {
|
||||
capture_critical_error(CriticalErrorCategory::Internal);
|
||||
42
|
||||
});
|
||||
|
||||
assert_eq!(product_result, 42);
|
||||
assert_eq!(transport.attempts.load(Ordering::Relaxed), 1);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,36 @@
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub enum OperationalIncident {
|
||||
InvocationHistoryLost,
|
||||
}
|
||||
|
||||
static INVOCATION_HISTORY_LOST_TOTAL: AtomicU64 = AtomicU64::new(0);
|
||||
|
||||
pub fn record_operational_incident(incident: OperationalIncident) {
|
||||
let counter = counter(incident);
|
||||
let _ = counter.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |value| {
|
||||
value.checked_add(1)
|
||||
});
|
||||
match incident {
|
||||
OperationalIncident::InvocationHistoryLost => {
|
||||
metrics::counter!("crank_invocation_history_lost_total").increment(1);
|
||||
metrics::counter!(
|
||||
"crank_telemetry_export_failures_total",
|
||||
"signal_type" => "invocation_history",
|
||||
"exporter" => "postgres"
|
||||
)
|
||||
.increment(1);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn operational_incident_total(incident: OperationalIncident) -> u64 {
|
||||
counter(incident).load(Ordering::Relaxed)
|
||||
}
|
||||
|
||||
fn counter(incident: OperationalIncident) -> &'static AtomicU64 {
|
||||
match incident {
|
||||
OperationalIncident::InvocationHistoryLost => &INVOCATION_HISTORY_LOST_TOTAL,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
use std::time::Instant;
|
||||
|
||||
use axum::{
|
||||
extract::{MatchedPath, Request},
|
||||
middleware::Next,
|
||||
response::Response,
|
||||
};
|
||||
use metrics::{Gauge, Unit};
|
||||
|
||||
use crate::{MetricKind, MetricUnit, metric_schema};
|
||||
|
||||
pub async fn record_http_request(request: Request, next: Next) -> Response {
|
||||
let route = request
|
||||
.extensions()
|
||||
.get::<MatchedPath>()
|
||||
.map_or("unmatched", MatchedPath::as_str)
|
||||
.to_owned();
|
||||
let method = normalized_http_method(request.method().as_str());
|
||||
let started_at = Instant::now();
|
||||
let _inflight = GaugeGuard::increment("crank_http_inflight");
|
||||
|
||||
let response = next.run(request).await;
|
||||
let status_class = status_class(response.status().as_u16());
|
||||
|
||||
metrics::counter!(
|
||||
"crank_http_requests_total",
|
||||
"route" => route.clone(),
|
||||
"method" => method,
|
||||
"status_class" => status_class
|
||||
)
|
||||
.increment(1);
|
||||
metrics::histogram!(
|
||||
"crank_http_request_duration_seconds",
|
||||
"route" => route,
|
||||
"method" => method
|
||||
)
|
||||
.record(started_at.elapsed().as_secs_f64());
|
||||
|
||||
response
|
||||
}
|
||||
|
||||
pub fn record_db_pool_connections(total: u32, idle: usize) {
|
||||
let idle = idle.min(total as usize) as f64;
|
||||
metrics::gauge!("crank_db_pool_connections", "state" => "idle").set(idle);
|
||||
metrics::gauge!("crank_db_pool_connections", "state" => "used").set(f64::from(total) - idle);
|
||||
}
|
||||
|
||||
pub(crate) fn register_metric_schema() {
|
||||
for definition in metric_schema() {
|
||||
let unit = match definition.unit {
|
||||
MetricUnit::Count => Unit::Count,
|
||||
MetricUnit::Seconds => Unit::Seconds,
|
||||
};
|
||||
match definition.kind {
|
||||
MetricKind::Counter => {
|
||||
metrics::describe_counter!(definition.name, unit, definition.description);
|
||||
}
|
||||
MetricKind::Gauge => {
|
||||
metrics::describe_gauge!(definition.name, unit, definition.description);
|
||||
}
|
||||
MetricKind::Histogram => {
|
||||
metrics::describe_histogram!(definition.name, unit, definition.description);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
metrics::gauge!("crank_http_inflight").set(0.0);
|
||||
metrics::gauge!("crank_mcp_active_sessions").set(0.0);
|
||||
metrics::gauge!("crank_runtime_inflight").set(0.0);
|
||||
metrics::gauge!("crank_db_pool_connections", "state" => "idle").set(0.0);
|
||||
metrics::gauge!("crank_db_pool_connections", "state" => "used").set(0.0);
|
||||
metrics::gauge!("crank_catalog_tools").set(0.0);
|
||||
metrics::gauge!("crank_catalog_estimated_context_tokens").set(0.0);
|
||||
metrics::gauge!("crank_catalog_warnings").set(0.0);
|
||||
}
|
||||
|
||||
fn normalized_http_method(method: &str) -> &'static str {
|
||||
match method {
|
||||
"GET" => "GET",
|
||||
"POST" => "POST",
|
||||
"PUT" => "PUT",
|
||||
"PATCH" => "PATCH",
|
||||
"DELETE" => "DELETE",
|
||||
"OPTIONS" => "OPTIONS",
|
||||
"HEAD" => "HEAD",
|
||||
"CONNECT" => "CONNECT",
|
||||
"TRACE" => "TRACE",
|
||||
_ => "OTHER",
|
||||
}
|
||||
}
|
||||
|
||||
fn status_class(status: u16) -> &'static str {
|
||||
match status {
|
||||
100..=199 => "1xx",
|
||||
200..=299 => "2xx",
|
||||
300..=399 => "3xx",
|
||||
400..=499 => "4xx",
|
||||
500..=599 => "5xx",
|
||||
_ => "other",
|
||||
}
|
||||
}
|
||||
|
||||
struct GaugeGuard {
|
||||
gauge: Gauge,
|
||||
}
|
||||
|
||||
impl GaugeGuard {
|
||||
fn increment(name: &'static str) -> Self {
|
||||
let gauge = metrics::gauge!(name);
|
||||
gauge.increment(1.0);
|
||||
Self { gauge }
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for GaugeGuard {
|
||||
fn drop(&mut self) {
|
||||
self.gauge.decrement(1.0);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{normalized_http_method, status_class};
|
||||
|
||||
#[test]
|
||||
fn normalizes_unbounded_http_values() {
|
||||
assert_eq!(normalized_http_method("GET"), "GET");
|
||||
assert_eq!(normalized_http_method("CUSTOM-user-controlled"), "OTHER");
|
||||
assert_eq!(status_class(204), "2xx");
|
||||
assert_eq!(status_class(429), "4xx");
|
||||
assert_eq!(status_class(999), "other");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
mod config;
|
||||
mod correlation;
|
||||
mod error_reporting;
|
||||
mod incidents;
|
||||
mod instrumentation;
|
||||
mod lifecycle;
|
||||
mod logging;
|
||||
mod metrics_schema;
|
||||
mod otlp;
|
||||
mod prometheus;
|
||||
mod propagation;
|
||||
mod redaction;
|
||||
mod schema;
|
||||
|
||||
pub use config::{ObservabilityConfig, ObservabilityConfigError, ServiceIdentity};
|
||||
pub use correlation::RequestId;
|
||||
pub use error_reporting::{
|
||||
CriticalErrorCategory, SentryConfig, SentryConfigError, capture_critical_error,
|
||||
with_request_correlation,
|
||||
};
|
||||
pub use incidents::{OperationalIncident, operational_incident_total, record_operational_incident};
|
||||
pub use instrumentation::{record_db_pool_connections, record_http_request};
|
||||
pub use lifecycle::{ObservabilityInitError, ObservabilityLifecycle, init};
|
||||
pub use logging::build_subscriber;
|
||||
pub use metrics_schema::{
|
||||
DURATION_BUCKETS_SECONDS, MetricDefinition, MetricKind, MetricUnit, metric_schema,
|
||||
};
|
||||
pub use otlp::{
|
||||
OtlpBatchConfig, OtlpTraceConfig, OtlpTraceConfigError, OtlpTraceError, build_tracer_provider,
|
||||
};
|
||||
pub use prometheus::{
|
||||
MetricsConfig, MetricsConfigError, MetricsServeError, MetricsSurface, MetricsSurfaceError,
|
||||
};
|
||||
pub use propagation::{inject_current_trace_context, set_remote_trace_parent};
|
||||
pub use redaction::{
|
||||
REDACTED_MARKER, RedactionLimits, RedactionLimitsError, SafeJsonError, redact_value, safe_json,
|
||||
};
|
||||
@@ -0,0 +1,90 @@
|
||||
use std::{fmt, io};
|
||||
|
||||
use thiserror::Error;
|
||||
use tracing_subscriber::util::SubscriberInitExt;
|
||||
|
||||
use crate::{
|
||||
MetricsConfig, MetricsSurface, MetricsSurfaceError, ObservabilityConfig,
|
||||
ObservabilityConfigError, OtlpTraceConfig, OtlpTraceConfigError, OtlpTraceError,
|
||||
RedactionLimitsError, SentryConfig, SentryConfigError, error_reporting::init_sentry,
|
||||
instrumentation::register_metric_schema, logging::build_subscriber_with_tracer,
|
||||
otlp::build_tracer_provider, prometheus::install_prometheus_recorder,
|
||||
propagation::install_trace_context_propagator,
|
||||
};
|
||||
|
||||
#[must_use = "observability resources must be retained until process shutdown"]
|
||||
pub struct ObservabilityLifecycle {
|
||||
metrics_handle: metrics_exporter_prometheus::PrometheusHandle,
|
||||
tracer_provider: Option<opentelemetry_sdk::trace::SdkTracerProvider>,
|
||||
sentry_guard: Option<sentry::ClientInitGuard>,
|
||||
}
|
||||
|
||||
impl ObservabilityLifecycle {
|
||||
pub fn init(config: ObservabilityConfig) -> Result<Self, ObservabilityInitError> {
|
||||
let identity = config.identity().clone();
|
||||
let redaction_limits = config.redaction_limits();
|
||||
let sentry_config = SentryConfig::from_env()?;
|
||||
let trace_config = OtlpTraceConfig::from_env()?;
|
||||
let tracing = build_tracer_provider(&identity, &trace_config)?;
|
||||
let tracer = tracing.as_ref().map(|(_, tracer)| tracer.clone());
|
||||
install_trace_context_propagator();
|
||||
build_subscriber_with_tracer(config, io::stdout, tracer)?
|
||||
.try_init()
|
||||
.map_err(|_| ObservabilityInitError::SubscriberAlreadyInitialized)?;
|
||||
let metrics_handle = install_prometheus_recorder(&identity)?;
|
||||
register_metric_schema();
|
||||
let sentry_guard = init_sentry(&identity, redaction_limits, sentry_config);
|
||||
|
||||
Ok(Self {
|
||||
metrics_handle,
|
||||
tracer_provider: tracing.map(|(provider, _)| provider),
|
||||
sentry_guard,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn metrics_surface(&self, config: MetricsConfig) -> MetricsSurface {
|
||||
MetricsSurface::new(config, self.metrics_handle.clone())
|
||||
}
|
||||
|
||||
pub fn traces_enabled(&self) -> bool {
|
||||
self.tracer_provider.is_some()
|
||||
}
|
||||
|
||||
pub fn critical_errors_enabled(&self) -> bool {
|
||||
self.sentry_guard.is_some()
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for ObservabilityLifecycle {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ObservabilityLifecycle")
|
||||
.field("traces_enabled", &self.traces_enabled())
|
||||
.field("critical_errors_enabled", &self.critical_errors_enabled())
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
pub fn init(config: ObservabilityConfig) -> Result<ObservabilityLifecycle, ObservabilityInitError> {
|
||||
ObservabilityLifecycle::init(config)
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum ObservabilityInitError {
|
||||
#[error(transparent)]
|
||||
InvalidConfig(#[from] ObservabilityConfigError),
|
||||
#[error(transparent)]
|
||||
InvalidRedactionLimits(#[from] RedactionLimitsError),
|
||||
#[error("invalid log filter")]
|
||||
InvalidFilter,
|
||||
#[error("global tracing subscriber is already initialized")]
|
||||
SubscriberAlreadyInitialized,
|
||||
#[error(transparent)]
|
||||
Metrics(#[from] MetricsSurfaceError),
|
||||
#[error(transparent)]
|
||||
OtlpConfig(#[from] OtlpTraceConfigError),
|
||||
#[error(transparent)]
|
||||
Otlp(#[from] OtlpTraceError),
|
||||
#[error(transparent)]
|
||||
SentryConfig(#[from] SentryConfigError),
|
||||
}
|
||||
@@ -0,0 +1,314 @@
|
||||
use std::fmt;
|
||||
|
||||
use serde_json::{Map, Number, Value};
|
||||
use time::{OffsetDateTime, format_description::well_known::Rfc3339};
|
||||
use tracing::{Event, Subscriber, field::Visit};
|
||||
use tracing_subscriber::{
|
||||
EnvFilter, Layer,
|
||||
filter::filter_fn,
|
||||
fmt::{FmtContext, FormatEvent, FormatFields, MakeWriter, format::Writer},
|
||||
layer::SubscriberExt,
|
||||
registry::LookupSpan,
|
||||
};
|
||||
|
||||
use crate::{
|
||||
ObservabilityConfig, ObservabilityInitError, RedactionLimits, ServiceIdentity,
|
||||
propagation::current_trace_id,
|
||||
redaction::{redact_value, truncate_string},
|
||||
schema::LogEnvelope,
|
||||
};
|
||||
|
||||
pub fn build_subscriber<W>(
|
||||
config: ObservabilityConfig,
|
||||
writer: W,
|
||||
) -> Result<impl Subscriber + Send + Sync, ObservabilityInitError>
|
||||
where
|
||||
W: for<'writer> MakeWriter<'writer> + Send + Sync + 'static,
|
||||
{
|
||||
build_subscriber_with_tracer(config, writer, None)
|
||||
}
|
||||
|
||||
pub(crate) fn build_subscriber_with_tracer<W>(
|
||||
config: ObservabilityConfig,
|
||||
writer: W,
|
||||
tracer: Option<opentelemetry_sdk::trace::SdkTracer>,
|
||||
) -> Result<impl Subscriber + Send + Sync, ObservabilityInitError>
|
||||
where
|
||||
W: for<'writer> MakeWriter<'writer> + Send + Sync + 'static,
|
||||
{
|
||||
let (identity, filter, limits) = config.into_parts();
|
||||
limits.validate()?;
|
||||
let filter = EnvFilter::try_new(filter).map_err(|_| ObservabilityInitError::InvalidFilter)?;
|
||||
let formatter = JsonEventFormatter::new(identity, limits);
|
||||
let fmt_layer = tracing_subscriber::fmt::layer()
|
||||
.with_ansi(false)
|
||||
.event_format(formatter)
|
||||
.with_writer(writer)
|
||||
.with_filter(filter);
|
||||
let otel_layer = tracer.map(|tracer| {
|
||||
tracing_opentelemetry::layer()
|
||||
.with_tracer(tracer)
|
||||
.with_filter(filter_fn(|metadata| {
|
||||
metadata.is_span() && metadata.target() == "crank::trace"
|
||||
}))
|
||||
});
|
||||
|
||||
Ok(tracing_subscriber::registry()
|
||||
.with(fmt_layer)
|
||||
.with(otel_layer))
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
struct JsonEventFormatter {
|
||||
identity: ServiceIdentity,
|
||||
limits: RedactionLimits,
|
||||
}
|
||||
|
||||
impl JsonEventFormatter {
|
||||
fn new(identity: ServiceIdentity, limits: RedactionLimits) -> Self {
|
||||
Self { identity, limits }
|
||||
}
|
||||
|
||||
fn envelope(&self, event: &Event<'_>) -> Result<LogEnvelope, fmt::Error> {
|
||||
let metadata = event.metadata();
|
||||
let mut visitor = JsonFieldVisitor::default();
|
||||
event.record(&mut visitor);
|
||||
let mut raw_fields = visitor.fields;
|
||||
let request_id = take_correlation_id(&mut raw_fields, "request_id")
|
||||
.map(|value| truncate_string(&value, self.limits.max_string_bytes));
|
||||
let trace_id = take_correlation_id(&mut raw_fields, "trace_id")
|
||||
.or_else(current_trace_id)
|
||||
.map(|value| truncate_string(&value, self.limits.max_string_bytes));
|
||||
let cleaned = redact_value(&Value::Object(raw_fields), self.limits);
|
||||
let fields = cleaned.as_object().cloned().unwrap_or_default();
|
||||
let timestamp = OffsetDateTime::now_utc()
|
||||
.format(&Rfc3339)
|
||||
.map_err(|_| fmt::Error)?;
|
||||
|
||||
Ok(LogEnvelope {
|
||||
timestamp,
|
||||
level: metadata.level().as_str().to_owned(),
|
||||
service: self.identity.service().to_owned(),
|
||||
version: self.identity.version().to_owned(),
|
||||
environment: self.identity.environment().to_owned(),
|
||||
target: truncate_string(metadata.target(), self.limits.max_string_bytes),
|
||||
event: truncate_string(metadata.name(), self.limits.max_string_bytes),
|
||||
request_id,
|
||||
trace_id,
|
||||
fields,
|
||||
})
|
||||
}
|
||||
|
||||
fn serialize_bounded(&self, mut envelope: LogEnvelope) -> Result<String, fmt::Error> {
|
||||
let line_budget = self.limits.max_event_bytes.saturating_sub(1);
|
||||
let serialized = serde_json::to_string(&envelope).map_err(|_| fmt::Error)?;
|
||||
if serialized.len() <= line_budget {
|
||||
return Ok(serialized);
|
||||
}
|
||||
|
||||
envelope.fields = Map::from_iter([("truncated".to_owned(), Value::Bool(true))]);
|
||||
let fallback_string_limit = self.limits.max_string_bytes.min(64);
|
||||
envelope.target = truncate_string(&envelope.target, fallback_string_limit);
|
||||
envelope.event = truncate_string(&envelope.event, fallback_string_limit);
|
||||
envelope.request_id = envelope
|
||||
.request_id
|
||||
.map(|value| truncate_string(&value, fallback_string_limit));
|
||||
envelope.trace_id = envelope
|
||||
.trace_id
|
||||
.map(|value| truncate_string(&value, fallback_string_limit));
|
||||
let serialized = serde_json::to_string(&envelope).map_err(|_| fmt::Error)?;
|
||||
if serialized.len() <= line_budget {
|
||||
return Ok(serialized);
|
||||
}
|
||||
|
||||
envelope.request_id = None;
|
||||
envelope.trace_id = None;
|
||||
envelope.target = truncate_string(&envelope.target, 16);
|
||||
envelope.event = truncate_string(&envelope.event, 16);
|
||||
let serialized = serde_json::to_string(&envelope).map_err(|_| fmt::Error)?;
|
||||
(serialized.len() <= line_budget)
|
||||
.then_some(serialized)
|
||||
.ok_or(fmt::Error)
|
||||
}
|
||||
}
|
||||
|
||||
impl<S, N> FormatEvent<S, N> for JsonEventFormatter
|
||||
where
|
||||
S: Subscriber + for<'lookup> LookupSpan<'lookup>,
|
||||
N: for<'writer> FormatFields<'writer> + 'static,
|
||||
{
|
||||
fn format_event(
|
||||
&self,
|
||||
_ctx: &FmtContext<'_, S, N>,
|
||||
mut writer: Writer<'_>,
|
||||
event: &Event<'_>,
|
||||
) -> fmt::Result {
|
||||
let serialized = self.serialize_bounded(self.envelope(event)?)?;
|
||||
writer.write_str(&serialized)?;
|
||||
writer.write_char('\n')
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct JsonFieldVisitor {
|
||||
fields: Map<String, Value>,
|
||||
}
|
||||
|
||||
impl JsonFieldVisitor {
|
||||
fn insert(&mut self, field: &tracing::field::Field, value: Value) {
|
||||
self.fields.insert(field.name().to_owned(), value);
|
||||
}
|
||||
}
|
||||
|
||||
impl Visit for JsonFieldVisitor {
|
||||
fn record_i64(&mut self, field: &tracing::field::Field, value: i64) {
|
||||
self.insert(field, Value::Number(value.into()));
|
||||
}
|
||||
|
||||
fn record_u64(&mut self, field: &tracing::field::Field, value: u64) {
|
||||
self.insert(field, Value::Number(value.into()));
|
||||
}
|
||||
|
||||
fn record_bool(&mut self, field: &tracing::field::Field, value: bool) {
|
||||
self.insert(field, Value::Bool(value));
|
||||
}
|
||||
|
||||
fn record_f64(&mut self, field: &tracing::field::Field, value: f64) {
|
||||
let value = Number::from_f64(value)
|
||||
.map(Value::Number)
|
||||
.unwrap_or(Value::Null);
|
||||
self.insert(field, value);
|
||||
}
|
||||
|
||||
fn record_str(&mut self, field: &tracing::field::Field, value: &str) {
|
||||
self.insert(field, Value::String(value.to_owned()));
|
||||
}
|
||||
|
||||
fn record_debug(&mut self, field: &tracing::field::Field, value: &dyn fmt::Debug) {
|
||||
let rendered = format!("{value:?}");
|
||||
let value = if is_correlation_field(field.name()) {
|
||||
Value::String(debug_scalar(&rendered))
|
||||
} else {
|
||||
match serde_json::from_str(&rendered) {
|
||||
Ok(value @ (Value::Object(_) | Value::Array(_))) => value,
|
||||
_ if field.name() == "message" || is_safe_display_scalar(&rendered) => {
|
||||
Value::String(rendered)
|
||||
}
|
||||
_ => Value::String(crate::REDACTED_MARKER.to_owned()),
|
||||
}
|
||||
};
|
||||
self.insert(field, value);
|
||||
}
|
||||
}
|
||||
|
||||
fn take_correlation_id(fields: &mut Map<String, Value>, name: &str) -> Option<String> {
|
||||
let value = fields.remove(name)?;
|
||||
let value = match value {
|
||||
Value::String(value) => value,
|
||||
Value::Number(value) => value.to_string(),
|
||||
Value::Bool(value) => value.to_string(),
|
||||
Value::Null | Value::Array(_) | Value::Object(_) => return None,
|
||||
};
|
||||
(!value.is_empty()).then_some(value)
|
||||
}
|
||||
|
||||
fn is_correlation_field(name: &str) -> bool {
|
||||
matches!(name, "request_id" | "trace_id" | "correlation_id")
|
||||
}
|
||||
|
||||
fn debug_scalar(rendered: &str) -> String {
|
||||
serde_json::from_str::<String>(rendered).unwrap_or_else(|_| rendered.to_owned())
|
||||
}
|
||||
|
||||
fn is_safe_display_scalar(rendered: &str) -> bool {
|
||||
!rendered.is_empty()
|
||||
&& rendered.bytes().all(|byte| {
|
||||
byte.is_ascii_alphanumeric()
|
||||
|| matches!(byte, b'-' | b'_' | b'.' | b':' | b'/' | b'+' | b'@')
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::{
|
||||
io,
|
||||
sync::{Arc, Mutex},
|
||||
};
|
||||
|
||||
use opentelemetry::{global, trace::TracerProvider as _};
|
||||
use opentelemetry_sdk::{
|
||||
error::OTelSdkResult,
|
||||
propagation::TraceContextPropagator,
|
||||
trace::{SdkTracerProvider, SpanData, SpanExporter},
|
||||
};
|
||||
use tracing::info;
|
||||
|
||||
use super::build_subscriber_with_tracer;
|
||||
use crate::{
|
||||
ObservabilityConfig, RedactionLimits, ServiceIdentity, inject_current_trace_context,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn trace_spans_ignore_the_log_level_filter() {
|
||||
global::set_text_map_propagator(TraceContextPropagator::new());
|
||||
let provider = SdkTracerProvider::builder().build();
|
||||
let tracer = provider.tracer("trace-filter-test");
|
||||
let subscriber = build_subscriber_with_tracer(test_config("warn"), io::sink, Some(tracer))
|
||||
.expect("subscriber must build");
|
||||
let dispatch = tracing::Dispatch::new(subscriber);
|
||||
let _dispatch_guard = tracing::dispatcher::set_default(&dispatch);
|
||||
let span = tracing::info_span!(target: "crank::trace", "http.request");
|
||||
let _span_guard = span.enter();
|
||||
let mut headers = axum::http::HeaderMap::new();
|
||||
|
||||
assert!(inject_current_trace_context(&mut headers));
|
||||
assert!(headers.contains_key("traceparent"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn otel_layer_does_not_export_events() {
|
||||
let exported = Arc::new(Mutex::new(Vec::new()));
|
||||
let provider = SdkTracerProvider::builder()
|
||||
.with_simple_exporter(CapturingExporter(Arc::clone(&exported)))
|
||||
.build();
|
||||
let tracer = provider.tracer("event-filter-test");
|
||||
let subscriber = build_subscriber_with_tracer(test_config("info"), io::sink, Some(tracer))
|
||||
.expect("subscriber must build");
|
||||
let dispatch = tracing::Dispatch::new(subscriber);
|
||||
|
||||
tracing::dispatcher::with_default(&dispatch, || {
|
||||
let span = tracing::info_span!(target: "crank::trace", "http.request");
|
||||
let _span_guard = span.enter();
|
||||
info!(password = "canary-secret", "sensitive event");
|
||||
});
|
||||
provider.force_flush().expect("span must be exported");
|
||||
|
||||
let spans = exported.lock().expect("capture lock");
|
||||
assert_eq!(spans.len(), 1);
|
||||
assert!(spans[0].events.is_empty());
|
||||
assert!(
|
||||
!format!("{:?}", spans[0])
|
||||
.as_bytes()
|
||||
.windows(b"canary-secret".len())
|
||||
.any(|window| window == b"canary-secret")
|
||||
);
|
||||
}
|
||||
|
||||
fn test_config(filter: &str) -> ObservabilityConfig {
|
||||
ObservabilityConfig::new(
|
||||
ServiceIdentity::try_new("admin-api", "test", "test").unwrap(),
|
||||
filter,
|
||||
RedactionLimits::default(),
|
||||
)
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
struct CapturingExporter(Arc<Mutex<Vec<SpanData>>>);
|
||||
|
||||
impl SpanExporter for CapturingExporter {
|
||||
async fn export(&self, batch: Vec<SpanData>) -> OTelSdkResult {
|
||||
self.0.lock().expect("capture lock").extend(batch);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,159 @@
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub enum MetricKind {
|
||||
Counter,
|
||||
Gauge,
|
||||
Histogram,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub enum MetricUnit {
|
||||
Count,
|
||||
Seconds,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub struct MetricDefinition {
|
||||
pub name: &'static str,
|
||||
pub kind: MetricKind,
|
||||
pub unit: MetricUnit,
|
||||
pub labels: &'static [&'static str],
|
||||
pub description: &'static str,
|
||||
}
|
||||
|
||||
pub const DURATION_BUCKETS_SECONDS: &[f64] = &[
|
||||
0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0, 30.0, 60.0,
|
||||
];
|
||||
|
||||
const METRIC_SCHEMA: &[MetricDefinition] = &[
|
||||
counter(
|
||||
"crank_http_requests_total",
|
||||
&["route", "method", "status_class"],
|
||||
"Total HTTP requests.",
|
||||
),
|
||||
histogram(
|
||||
"crank_http_request_duration_seconds",
|
||||
&["route", "method"],
|
||||
"HTTP request duration in seconds.",
|
||||
),
|
||||
gauge(
|
||||
"crank_http_inflight",
|
||||
&[],
|
||||
"HTTP requests currently being processed.",
|
||||
),
|
||||
counter(
|
||||
"crank_mcp_requests_total",
|
||||
&["method", "response_mode", "outcome"],
|
||||
"Total MCP JSON-RPC requests.",
|
||||
),
|
||||
gauge(
|
||||
"crank_mcp_active_sessions",
|
||||
&[],
|
||||
"Active MCP transport sessions.",
|
||||
),
|
||||
counter(
|
||||
"crank_tool_invocations_total",
|
||||
&["source", "outcome", "error_kind"],
|
||||
"Total tool invocations.",
|
||||
),
|
||||
histogram(
|
||||
"crank_tool_invocation_duration_seconds",
|
||||
&["source", "outcome"],
|
||||
"Tool invocation duration in seconds.",
|
||||
),
|
||||
counter(
|
||||
"crank_upstream_requests_total",
|
||||
&["operation_kind", "outcome"],
|
||||
"Total upstream requests.",
|
||||
),
|
||||
histogram(
|
||||
"crank_upstream_request_duration_seconds",
|
||||
&["operation_kind", "outcome"],
|
||||
"Upstream request duration in seconds.",
|
||||
),
|
||||
gauge(
|
||||
"crank_runtime_inflight",
|
||||
&[],
|
||||
"Runtime executions currently in progress.",
|
||||
),
|
||||
counter(
|
||||
"crank_runtime_limit_rejections_total",
|
||||
&["stage"],
|
||||
"Runtime executions rejected by a bounded limit.",
|
||||
),
|
||||
gauge(
|
||||
"crank_db_pool_connections",
|
||||
&["state"],
|
||||
"PostgreSQL pool connections by state.",
|
||||
),
|
||||
gauge(
|
||||
"crank_catalog_tools",
|
||||
&[],
|
||||
"Tools in the current published catalog.",
|
||||
),
|
||||
gauge(
|
||||
"crank_catalog_estimated_context_tokens",
|
||||
&[],
|
||||
"Estimated context tokens in the current published catalog.",
|
||||
),
|
||||
gauge(
|
||||
"crank_catalog_warnings",
|
||||
&[],
|
||||
"Warnings in the current published catalog.",
|
||||
),
|
||||
counter(
|
||||
"crank_invocation_history_lost_total",
|
||||
&[],
|
||||
"Invocation history records lost after an action completed.",
|
||||
),
|
||||
counter(
|
||||
"crank_telemetry_export_failures_total",
|
||||
&["signal_type", "exporter"],
|
||||
"Telemetry export failures.",
|
||||
),
|
||||
];
|
||||
|
||||
pub const fn metric_schema() -> &'static [MetricDefinition] {
|
||||
METRIC_SCHEMA
|
||||
}
|
||||
|
||||
const fn counter(
|
||||
name: &'static str,
|
||||
labels: &'static [&'static str],
|
||||
description: &'static str,
|
||||
) -> MetricDefinition {
|
||||
MetricDefinition {
|
||||
name,
|
||||
kind: MetricKind::Counter,
|
||||
unit: MetricUnit::Count,
|
||||
labels,
|
||||
description,
|
||||
}
|
||||
}
|
||||
|
||||
const fn gauge(
|
||||
name: &'static str,
|
||||
labels: &'static [&'static str],
|
||||
description: &'static str,
|
||||
) -> MetricDefinition {
|
||||
MetricDefinition {
|
||||
name,
|
||||
kind: MetricKind::Gauge,
|
||||
unit: MetricUnit::Count,
|
||||
labels,
|
||||
description,
|
||||
}
|
||||
}
|
||||
|
||||
const fn histogram(
|
||||
name: &'static str,
|
||||
labels: &'static [&'static str],
|
||||
description: &'static str,
|
||||
) -> MetricDefinition {
|
||||
MetricDefinition {
|
||||
name,
|
||||
kind: MetricKind::Histogram,
|
||||
unit: MetricUnit::Seconds,
|
||||
labels,
|
||||
description,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,954 @@
|
||||
use std::{collections::HashMap, env, fmt, time::Duration};
|
||||
|
||||
use axum::http::{HeaderName, HeaderValue};
|
||||
use opentelemetry::{
|
||||
KeyValue, Value,
|
||||
trace::{Status, TracerProvider as _},
|
||||
};
|
||||
use opentelemetry_otlp::{Protocol, SpanExporter, WithExportConfig, WithHttpConfig};
|
||||
use opentelemetry_sdk::{
|
||||
Resource,
|
||||
error::OTelSdkResult,
|
||||
trace::{
|
||||
BatchConfigBuilder, BatchSpanProcessor, SdkTracer, SdkTracerProvider, SpanData,
|
||||
SpanExporter as SpanExporterTrait,
|
||||
},
|
||||
};
|
||||
use percent_encoding::percent_decode_str;
|
||||
use thiserror::Error;
|
||||
use url::Url;
|
||||
|
||||
use crate::ServiceIdentity;
|
||||
|
||||
const DEFAULT_EXPORT_TIMEOUT: Duration = Duration::from_secs(10);
|
||||
const DEFAULT_MAX_QUEUE_SIZE: usize = 2_048;
|
||||
const DEFAULT_MAX_EXPORT_BATCH_SIZE: usize = 512;
|
||||
const DEFAULT_SCHEDULE_DELAY: Duration = Duration::from_secs(5);
|
||||
const DEFAULT_BATCH_EXPORT_TIMEOUT: Duration = Duration::from_secs(30);
|
||||
const MAX_QUEUE_SIZE: usize = 65_536;
|
||||
const MAX_DURATION: Duration = Duration::from_secs(300);
|
||||
|
||||
#[derive(Clone, Debug, Eq, PartialEq)]
|
||||
pub struct OtlpBatchConfig {
|
||||
max_queue_size: usize,
|
||||
max_export_batch_size: usize,
|
||||
scheduled_delay: Duration,
|
||||
export_timeout: Duration,
|
||||
}
|
||||
|
||||
impl OtlpBatchConfig {
|
||||
pub fn try_new(
|
||||
max_queue_size: usize,
|
||||
max_export_batch_size: usize,
|
||||
scheduled_delay: Duration,
|
||||
export_timeout: Duration,
|
||||
) -> Result<Self, OtlpTraceConfigError> {
|
||||
let valid = max_queue_size > 0
|
||||
&& max_queue_size <= MAX_QUEUE_SIZE
|
||||
&& max_export_batch_size > 0
|
||||
&& max_export_batch_size <= max_queue_size
|
||||
&& duration_is_bounded(scheduled_delay)
|
||||
&& duration_is_bounded(export_timeout);
|
||||
if !valid {
|
||||
return Err(OtlpTraceConfigError::InvalidBatchLimits);
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
max_queue_size,
|
||||
max_export_batch_size,
|
||||
scheduled_delay,
|
||||
export_timeout,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn max_queue_size(&self) -> usize {
|
||||
self.max_queue_size
|
||||
}
|
||||
|
||||
pub fn max_export_batch_size(&self) -> usize {
|
||||
self.max_export_batch_size
|
||||
}
|
||||
|
||||
pub fn scheduled_delay(&self) -> Duration {
|
||||
self.scheduled_delay
|
||||
}
|
||||
|
||||
pub fn export_timeout(&self) -> Duration {
|
||||
self.export_timeout
|
||||
}
|
||||
|
||||
fn sdk_config(&self) -> opentelemetry_sdk::trace::BatchConfig {
|
||||
BatchConfigBuilder::default()
|
||||
.with_max_queue_size(self.max_queue_size)
|
||||
.with_max_export_batch_size(self.max_export_batch_size)
|
||||
.with_scheduled_delay(self.scheduled_delay)
|
||||
.build()
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for OtlpBatchConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
max_queue_size: DEFAULT_MAX_QUEUE_SIZE,
|
||||
max_export_batch_size: DEFAULT_MAX_EXPORT_BATCH_SIZE,
|
||||
scheduled_delay: DEFAULT_SCHEDULE_DELAY,
|
||||
export_timeout: DEFAULT_BATCH_EXPORT_TIMEOUT,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Eq, PartialEq)]
|
||||
pub struct OtlpTraceConfig {
|
||||
endpoint: Option<String>,
|
||||
export_timeout: Duration,
|
||||
batch: OtlpBatchConfig,
|
||||
headers: HashMap<String, String>,
|
||||
}
|
||||
|
||||
impl fmt::Debug for OtlpTraceConfig {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("OtlpTraceConfig")
|
||||
.field("enabled", &self.is_enabled())
|
||||
.field("export_timeout", &self.export_timeout)
|
||||
.field("batch", &self.batch)
|
||||
.field("header_count", &self.headers.len())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl OtlpTraceConfig {
|
||||
pub fn from_env() -> Result<Self, OtlpTraceConfigError> {
|
||||
OtlpEnvSettings::from_env()?.into_config()
|
||||
}
|
||||
|
||||
fn from_settings(settings: OtlpEnvSettings) -> Result<Self, OtlpTraceConfigError> {
|
||||
let endpoint = match settings.traces_endpoint {
|
||||
Some(endpoint) => Some(validate_endpoint(endpoint, EndpointKind::Trace)?),
|
||||
None => settings
|
||||
.generic_endpoint
|
||||
.map(|endpoint| validate_endpoint(endpoint, EndpointKind::Generic))
|
||||
.transpose()?,
|
||||
};
|
||||
if endpoint.is_none() {
|
||||
return Ok(Self {
|
||||
endpoint: None,
|
||||
export_timeout: DEFAULT_EXPORT_TIMEOUT,
|
||||
batch: OtlpBatchConfig::default(),
|
||||
headers: HashMap::new(),
|
||||
});
|
||||
}
|
||||
|
||||
let protocol = settings.traces_protocol.or(settings.generic_protocol);
|
||||
let export_timeout = match settings.traces_timeout {
|
||||
Some(timeout) => duration_env("OTEL_EXPORTER_OTLP_TRACES_TIMEOUT", Some(timeout))?,
|
||||
None => duration_env("OTEL_EXPORTER_OTLP_TIMEOUT", settings.generic_timeout)?,
|
||||
}
|
||||
.unwrap_or(DEFAULT_EXPORT_TIMEOUT);
|
||||
let batch = OtlpBatchConfig::try_new(
|
||||
usize_env("OTEL_BSP_MAX_QUEUE_SIZE", settings.max_queue_size)?
|
||||
.unwrap_or(DEFAULT_MAX_QUEUE_SIZE),
|
||||
usize_env(
|
||||
"OTEL_BSP_MAX_EXPORT_BATCH_SIZE",
|
||||
settings.max_export_batch_size,
|
||||
)?
|
||||
.unwrap_or(DEFAULT_MAX_EXPORT_BATCH_SIZE),
|
||||
duration_env("OTEL_BSP_SCHEDULE_DELAY", settings.scheduled_delay)?
|
||||
.unwrap_or(DEFAULT_SCHEDULE_DELAY),
|
||||
duration_env("OTEL_BSP_EXPORT_TIMEOUT", settings.batch_export_timeout)?
|
||||
.unwrap_or(DEFAULT_BATCH_EXPORT_TIMEOUT),
|
||||
)?;
|
||||
let headers = settings
|
||||
.traces_headers
|
||||
.filter(|value| !value.is_empty())
|
||||
.or(settings.generic_headers.filter(|value| !value.is_empty()))
|
||||
.map(|value| parse_headers(&value))
|
||||
.transpose()?
|
||||
.unwrap_or_default();
|
||||
|
||||
Self::try_new_with_headers(endpoint, protocol, export_timeout, batch, headers)
|
||||
}
|
||||
|
||||
pub fn try_new(
|
||||
endpoint: Option<String>,
|
||||
protocol: Option<String>,
|
||||
export_timeout: Duration,
|
||||
batch: OtlpBatchConfig,
|
||||
) -> Result<Self, OtlpTraceConfigError> {
|
||||
Self::try_new_with_headers(endpoint, protocol, export_timeout, batch, HashMap::new())
|
||||
}
|
||||
|
||||
fn try_new_with_headers(
|
||||
endpoint: Option<String>,
|
||||
protocol: Option<String>,
|
||||
export_timeout: Duration,
|
||||
batch: OtlpBatchConfig,
|
||||
headers: HashMap<String, String>,
|
||||
) -> Result<Self, OtlpTraceConfigError> {
|
||||
let endpoint = endpoint
|
||||
.map(|endpoint| validate_endpoint(endpoint, EndpointKind::Trace))
|
||||
.transpose()?;
|
||||
if endpoint.is_some() && protocol.as_deref().unwrap_or("http/protobuf") != "http/protobuf" {
|
||||
return Err(OtlpTraceConfigError::UnsupportedProtocol);
|
||||
}
|
||||
if !duration_is_bounded(export_timeout) {
|
||||
return Err(OtlpTraceConfigError::InvalidDuration {
|
||||
field: "OTEL_EXPORTER_OTLP_TIMEOUT",
|
||||
});
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
endpoint,
|
||||
export_timeout,
|
||||
batch,
|
||||
headers,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn is_enabled(&self) -> bool {
|
||||
self.endpoint.is_some()
|
||||
}
|
||||
|
||||
pub fn export_timeout(&self) -> Duration {
|
||||
self.export_timeout
|
||||
}
|
||||
|
||||
pub fn batch(&self) -> &OtlpBatchConfig {
|
||||
&self.batch
|
||||
}
|
||||
|
||||
fn effective_export_timeout(&self) -> Duration {
|
||||
self.export_timeout.min(self.batch.export_timeout)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
fn endpoint(&self) -> Option<&str> {
|
||||
self.endpoint.as_deref()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
fn header(&self, name: &str) -> Option<&str> {
|
||||
self.headers.get(name).map(String::as_str)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct OtlpEnvSettings {
|
||||
traces_endpoint: Option<String>,
|
||||
generic_endpoint: Option<String>,
|
||||
traces_protocol: Option<String>,
|
||||
generic_protocol: Option<String>,
|
||||
traces_timeout: Option<String>,
|
||||
generic_timeout: Option<String>,
|
||||
traces_headers: Option<String>,
|
||||
generic_headers: Option<String>,
|
||||
max_queue_size: Option<String>,
|
||||
max_export_batch_size: Option<String>,
|
||||
scheduled_delay: Option<String>,
|
||||
batch_export_timeout: Option<String>,
|
||||
}
|
||||
|
||||
impl OtlpEnvSettings {
|
||||
fn from_env() -> Result<Self, OtlpTraceConfigError> {
|
||||
Ok(Self {
|
||||
traces_endpoint: optional_env("OTEL_EXPORTER_OTLP_TRACES_ENDPOINT")?,
|
||||
generic_endpoint: optional_env("OTEL_EXPORTER_OTLP_ENDPOINT")?,
|
||||
traces_protocol: optional_env("OTEL_EXPORTER_OTLP_TRACES_PROTOCOL")?,
|
||||
generic_protocol: optional_env("OTEL_EXPORTER_OTLP_PROTOCOL")?,
|
||||
traces_timeout: optional_env("OTEL_EXPORTER_OTLP_TRACES_TIMEOUT")?,
|
||||
generic_timeout: optional_env("OTEL_EXPORTER_OTLP_TIMEOUT")?,
|
||||
traces_headers: optional_env("OTEL_EXPORTER_OTLP_TRACES_HEADERS")?,
|
||||
generic_headers: optional_env("OTEL_EXPORTER_OTLP_HEADERS")?,
|
||||
max_queue_size: optional_env("OTEL_BSP_MAX_QUEUE_SIZE")?,
|
||||
max_export_batch_size: optional_env("OTEL_BSP_MAX_EXPORT_BATCH_SIZE")?,
|
||||
scheduled_delay: optional_env("OTEL_BSP_SCHEDULE_DELAY")?,
|
||||
batch_export_timeout: optional_env("OTEL_BSP_EXPORT_TIMEOUT")?,
|
||||
})
|
||||
}
|
||||
|
||||
fn into_config(self) -> Result<OtlpTraceConfig, OtlpTraceConfigError> {
|
||||
OtlpTraceConfig::from_settings(self)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Error, Eq, PartialEq)]
|
||||
pub enum OtlpTraceConfigError {
|
||||
#[error("OTLP environment variable is not valid UTF-8: {field}")]
|
||||
InvalidEnvironmentEncoding { field: &'static str },
|
||||
#[error("OTLP trace endpoint is invalid: {reason}")]
|
||||
InvalidEndpoint { reason: &'static str },
|
||||
#[error("OTLP trace protocol must be http/protobuf")]
|
||||
UnsupportedProtocol,
|
||||
#[error("OTLP numeric setting is invalid: {field}")]
|
||||
InvalidNumber { field: &'static str },
|
||||
#[error("OTLP duration setting is invalid: {field}")]
|
||||
InvalidDuration { field: &'static str },
|
||||
#[error("OTLP batch limits are invalid")]
|
||||
InvalidBatchLimits,
|
||||
#[error("OTLP trace headers are invalid")]
|
||||
InvalidHeaders,
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum OtlpTraceError {
|
||||
#[error("failed to configure OTLP trace exporter")]
|
||||
ExporterConfiguration,
|
||||
}
|
||||
|
||||
pub fn build_tracer_provider(
|
||||
identity: &ServiceIdentity,
|
||||
config: &OtlpTraceConfig,
|
||||
) -> Result<Option<(SdkTracerProvider, SdkTracer)>, OtlpTraceError> {
|
||||
let Some(endpoint) = config.endpoint.as_deref() else {
|
||||
return Ok(None);
|
||||
};
|
||||
let exporter = SpanExporter::builder()
|
||||
.with_http()
|
||||
.with_protocol(Protocol::HttpBinary)
|
||||
.with_endpoint(endpoint)
|
||||
.with_timeout(config.effective_export_timeout())
|
||||
.with_headers(config.headers.clone())
|
||||
.build()
|
||||
.map_err(|_| OtlpTraceError::ExporterConfiguration)?;
|
||||
let processor = BatchSpanProcessor::builder(ObservedSpanExporter(exporter))
|
||||
.with_batch_config(config.batch.sdk_config())
|
||||
.build();
|
||||
let resource = Resource::builder_empty()
|
||||
.with_attributes([
|
||||
KeyValue::new("service.name", identity.service().to_owned()),
|
||||
KeyValue::new("service.version", identity.version().to_owned()),
|
||||
KeyValue::new(
|
||||
"deployment.environment.name",
|
||||
identity.environment().to_owned(),
|
||||
),
|
||||
])
|
||||
.build();
|
||||
let provider = SdkTracerProvider::builder()
|
||||
.with_span_processor(processor)
|
||||
.with_resource(resource)
|
||||
.build();
|
||||
let tracer = provider.tracer("crank");
|
||||
|
||||
Ok(Some((provider, tracer)))
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct ObservedSpanExporter(SpanExporter);
|
||||
|
||||
impl SpanExporterTrait for ObservedSpanExporter {
|
||||
async fn export(&self, mut batch: Vec<SpanData>) -> OTelSdkResult {
|
||||
sanitize_trace_batch(&mut batch);
|
||||
let result = self.0.export(batch).await;
|
||||
if result.is_err() {
|
||||
metrics::counter!(
|
||||
"crank_telemetry_export_failures_total",
|
||||
"signal_type" => "trace",
|
||||
"exporter" => "otlp"
|
||||
)
|
||||
.increment(1);
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
fn shutdown_with_timeout(&self, timeout: Duration) -> OTelSdkResult {
|
||||
self.0.shutdown_with_timeout(timeout)
|
||||
}
|
||||
|
||||
fn force_flush(&self) -> OTelSdkResult {
|
||||
self.0.force_flush()
|
||||
}
|
||||
|
||||
fn set_resource(&mut self, resource: &Resource) {
|
||||
self.0.set_resource(resource);
|
||||
}
|
||||
}
|
||||
|
||||
fn sanitize_trace_batch(batch: &mut Vec<SpanData>) {
|
||||
batch.retain(|span| is_allowed_span_name(span.name.as_ref()));
|
||||
for span in batch {
|
||||
let original_attribute_count = span.attributes.len();
|
||||
span.attributes.retain(is_allowed_span_attribute);
|
||||
span.dropped_attributes_count = span
|
||||
.dropped_attributes_count
|
||||
.saturating_add((original_attribute_count - span.attributes.len()) as u32);
|
||||
span.events = Default::default();
|
||||
span.links = Default::default();
|
||||
if matches!(span.status, Status::Error { .. }) {
|
||||
span.status = Status::error("");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn is_allowed_span_name(name: &str) -> bool {
|
||||
matches!(
|
||||
name,
|
||||
"http.request"
|
||||
| "mcp.request"
|
||||
| "mcp.rate_limit"
|
||||
| "mcp.access.check"
|
||||
| "mcp.catalog.load"
|
||||
| "mcp.tools.resolve"
|
||||
| "approval.check"
|
||||
| "runtime.execute"
|
||||
| "runtime.arguments.map"
|
||||
| "runtime.idempotency"
|
||||
| "upstream.http"
|
||||
| "runtime.response.transform"
|
||||
| "auth.resolve"
|
||||
| "approval.recovery"
|
||||
| "history.write"
|
||||
| "db.query"
|
||||
)
|
||||
}
|
||||
|
||||
fn is_allowed_span_attribute(attribute: &KeyValue) -> bool {
|
||||
let Value::String(value) = &attribute.value else {
|
||||
return false;
|
||||
};
|
||||
let value = value.as_str();
|
||||
match attribute.key.as_str() {
|
||||
"request_id" => crate::RequestId::is_valid(value),
|
||||
"stage" => is_allowed_span_name(value),
|
||||
"outcome" => matches!(
|
||||
value,
|
||||
"success"
|
||||
| "error"
|
||||
| "allowed"
|
||||
| "denied"
|
||||
| "required"
|
||||
| "replay"
|
||||
| "execute"
|
||||
| "skipped"
|
||||
| "cache_hit"
|
||||
),
|
||||
"error.category" => matches!(
|
||||
value,
|
||||
"access"
|
||||
| "rate_limit"
|
||||
| "catalog"
|
||||
| "approval"
|
||||
| "idempotency"
|
||||
| "schema"
|
||||
| "mapping"
|
||||
| "upstream"
|
||||
| "transformation"
|
||||
| "history"
|
||||
| "database"
|
||||
| "concurrency"
|
||||
| "configuration"
|
||||
| "internal"
|
||||
),
|
||||
"db.system" => value == "postgresql",
|
||||
"db.operation" => matches!(
|
||||
value,
|
||||
"machine_access.read"
|
||||
| "machine_access.touch"
|
||||
| "catalog.load"
|
||||
| "approval.read"
|
||||
| "approval.write"
|
||||
| "auth_profile.read"
|
||||
| "secret.read"
|
||||
| "secret.touch"
|
||||
| "invocation_history.write"
|
||||
),
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum EndpointKind {
|
||||
Trace,
|
||||
Generic,
|
||||
}
|
||||
|
||||
fn validate_endpoint(endpoint: String, kind: EndpointKind) -> Result<String, OtlpTraceConfigError> {
|
||||
let mut url = Url::parse(&endpoint).map_err(|_| OtlpTraceConfigError::InvalidEndpoint {
|
||||
reason: "invalid URL",
|
||||
})?;
|
||||
if !matches!(url.scheme(), "http" | "https") {
|
||||
return Err(OtlpTraceConfigError::InvalidEndpoint {
|
||||
reason: "unsupported scheme",
|
||||
});
|
||||
}
|
||||
if url.host_str().is_none() {
|
||||
return Err(OtlpTraceConfigError::InvalidEndpoint {
|
||||
reason: "host is required",
|
||||
});
|
||||
}
|
||||
if !url.username().is_empty() || url.password().is_some() {
|
||||
return Err(OtlpTraceConfigError::InvalidEndpoint {
|
||||
reason: "credentials are forbidden",
|
||||
});
|
||||
}
|
||||
if url.query().is_some() || url.fragment().is_some() {
|
||||
return Err(OtlpTraceConfigError::InvalidEndpoint {
|
||||
reason: "query and fragment are forbidden",
|
||||
});
|
||||
}
|
||||
if matches!(kind, EndpointKind::Generic) {
|
||||
let path = url.path().trim_end_matches('/');
|
||||
url.set_path(&format!("{path}/v1/traces"));
|
||||
}
|
||||
|
||||
Ok(url.into())
|
||||
}
|
||||
|
||||
fn parse_headers(value: &str) -> Result<HashMap<String, String>, OtlpTraceConfigError> {
|
||||
value
|
||||
.split_terminator(',')
|
||||
.map(str::trim)
|
||||
.filter(|item| !item.is_empty())
|
||||
.try_fold(HashMap::new(), |mut headers, item| {
|
||||
let (name, encoded_value) = item
|
||||
.split_once('=')
|
||||
.ok_or(OtlpTraceConfigError::InvalidHeaders)?;
|
||||
let name = HeaderName::from_bytes(name.trim().as_bytes())
|
||||
.map_err(|_| OtlpTraceConfigError::InvalidHeaders)?;
|
||||
let value = percent_decode_str(encoded_value.trim())
|
||||
.decode_utf8()
|
||||
.map_err(|_| OtlpTraceConfigError::InvalidHeaders)?
|
||||
.into_owned();
|
||||
if value.is_empty() || HeaderValue::from_str(&value).is_err() {
|
||||
return Err(OtlpTraceConfigError::InvalidHeaders);
|
||||
}
|
||||
headers.insert(name.as_str().to_owned(), value);
|
||||
Ok(headers)
|
||||
})
|
||||
}
|
||||
|
||||
fn optional_env(field: &'static str) -> Result<Option<String>, OtlpTraceConfigError> {
|
||||
match env::var(field) {
|
||||
Ok(value) if value.is_empty() => Ok(None),
|
||||
Ok(value) => Ok(Some(value)),
|
||||
Err(env::VarError::NotPresent) => Ok(None),
|
||||
Err(env::VarError::NotUnicode(_)) => {
|
||||
Err(OtlpTraceConfigError::InvalidEnvironmentEncoding { field })
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn usize_env(
|
||||
field: &'static str,
|
||||
value: Option<String>,
|
||||
) -> Result<Option<usize>, OtlpTraceConfigError> {
|
||||
value
|
||||
.map(|value| {
|
||||
value
|
||||
.parse()
|
||||
.map_err(|_| OtlpTraceConfigError::InvalidNumber { field })
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn duration_env(
|
||||
field: &'static str,
|
||||
value: Option<String>,
|
||||
) -> Result<Option<Duration>, OtlpTraceConfigError> {
|
||||
value
|
||||
.map(|value| {
|
||||
value
|
||||
.parse::<u64>()
|
||||
.map(Duration::from_millis)
|
||||
.map_err(|_| OtlpTraceConfigError::InvalidDuration { field })
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn duration_is_bounded(duration: Duration) -> bool {
|
||||
!duration.is_zero() && duration <= MAX_DURATION
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::{
|
||||
io::{Read, Write},
|
||||
net::TcpListener,
|
||||
sync::mpsc,
|
||||
thread,
|
||||
time::{Duration, Instant},
|
||||
};
|
||||
|
||||
use axum::{
|
||||
Router,
|
||||
body::{Body, to_bytes},
|
||||
extract::Request,
|
||||
middleware::Next,
|
||||
response::Response,
|
||||
routing::get,
|
||||
};
|
||||
use opentelemetry::{
|
||||
KeyValue,
|
||||
trace::{Span as _, Status, Tracer as _},
|
||||
};
|
||||
use opentelemetry_proto::tonic::{
|
||||
collector::trace::v1::ExportTraceServiceRequest, common::v1::any_value,
|
||||
};
|
||||
use prost::Message;
|
||||
use tower::ServiceExt;
|
||||
use tracing::{Instrument, info_span};
|
||||
use tracing_subscriber::layer::SubscriberExt;
|
||||
|
||||
use super::{OtlpBatchConfig, OtlpEnvSettings, OtlpTraceConfig, build_tracer_provider};
|
||||
use crate::ServiceIdentity;
|
||||
|
||||
#[test]
|
||||
fn signal_specific_settings_override_generic_settings() {
|
||||
let config = OtlpEnvSettings {
|
||||
traces_endpoint: Some("https://traces.example.test/custom".to_owned()),
|
||||
generic_endpoint: Some("https://generic.example.test/otel".to_owned()),
|
||||
traces_protocol: Some("http/protobuf".to_owned()),
|
||||
generic_protocol: Some("grpc".to_owned()),
|
||||
traces_timeout: Some("2500".to_owned()),
|
||||
generic_timeout: Some("invalid-unused-fallback".to_owned()),
|
||||
..OtlpEnvSettings::default()
|
||||
}
|
||||
.into_config()
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
config.endpoint(),
|
||||
Some("https://traces.example.test/custom")
|
||||
);
|
||||
assert_eq!(config.export_timeout(), Duration::from_millis(2500));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn disabled_export_ignores_inactive_settings() {
|
||||
let config = OtlpEnvSettings {
|
||||
traces_protocol: Some("grpc".to_owned()),
|
||||
generic_protocol: Some("grpc".to_owned()),
|
||||
traces_timeout: Some("invalid".to_owned()),
|
||||
generic_timeout: Some("invalid".to_owned()),
|
||||
max_queue_size: Some("invalid".to_owned()),
|
||||
max_export_batch_size: Some("invalid".to_owned()),
|
||||
scheduled_delay: Some("invalid".to_owned()),
|
||||
batch_export_timeout: Some("invalid".to_owned()),
|
||||
..OtlpEnvSettings::default()
|
||||
}
|
||||
.into_config()
|
||||
.unwrap();
|
||||
|
||||
assert!(!config.is_enabled());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn empty_signal_headers_use_generic_headers() {
|
||||
let config = OtlpEnvSettings {
|
||||
traces_endpoint: Some("https://traces.example.test/v1/traces".to_owned()),
|
||||
traces_headers: Some(String::new()),
|
||||
generic_headers: Some(
|
||||
"authorization=Bearer%20canary-token,x-tenant=community".to_owned(),
|
||||
),
|
||||
..OtlpEnvSettings::default()
|
||||
}
|
||||
.into_config()
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(config.header("authorization"), Some("Bearer canary-token"));
|
||||
assert_eq!(config.header("x-tenant"), Some("community"));
|
||||
assert!(!format!("{config:?}").contains("canary-token"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_headers_return_a_safe_error() {
|
||||
let config = OtlpEnvSettings {
|
||||
traces_endpoint: Some("https://traces.example.test/v1/traces".to_owned()),
|
||||
traces_headers: Some("authorization=canary-secret%0Ainjected".to_owned()),
|
||||
..OtlpEnvSettings::default()
|
||||
};
|
||||
|
||||
let error = config.into_config().unwrap_err();
|
||||
assert!(matches!(error, super::OtlpTraceConfigError::InvalidHeaders));
|
||||
assert!(!error.to_string().contains("canary-secret"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stricter_batch_timeout_bounds_http_export() {
|
||||
let config = OtlpEnvSettings {
|
||||
traces_endpoint: Some("https://traces.example.test/v1/traces".to_owned()),
|
||||
traces_protocol: Some("http/protobuf".to_owned()),
|
||||
traces_timeout: Some("9000".to_owned()),
|
||||
batch_export_timeout: Some("2500".to_owned()),
|
||||
..OtlpEnvSettings::default()
|
||||
}
|
||||
.into_config()
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
config.effective_export_timeout(),
|
||||
Duration::from_millis(2500)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generic_endpoint_receives_standard_trace_path() {
|
||||
let config = OtlpEnvSettings {
|
||||
generic_endpoint: Some("https://generic.example.test/otel/".to_owned()),
|
||||
generic_protocol: Some("http/protobuf".to_owned()),
|
||||
..OtlpEnvSettings::default()
|
||||
}
|
||||
.into_config()
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
config.endpoint(),
|
||||
Some("https://generic.example.test/otel/v1/traces")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn disabled_export_does_not_build_a_provider() {
|
||||
let config = OtlpTraceConfig::try_new(
|
||||
None,
|
||||
None,
|
||||
Duration::from_secs(1),
|
||||
OtlpBatchConfig::default(),
|
||||
)
|
||||
.unwrap();
|
||||
let identity = ServiceIdentity::try_new("admin-api", "0.3.1", "test").unwrap();
|
||||
|
||||
assert!(build_tracer_provider(&identity, &config).unwrap().is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn real_http_protobuf_export_contains_resource_and_trace() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let (request_tx, request_rx) = mpsc::sync_channel(1);
|
||||
let server = thread::spawn(move || {
|
||||
let (mut stream, _) = listener.accept().unwrap();
|
||||
let request = read_http_request(&mut stream);
|
||||
stream
|
||||
.write_all(
|
||||
b"HTTP/1.1 200 OK\r\ncontent-type: application/x-protobuf\r\ncontent-length: 0\r\nconnection: close\r\n\r\n",
|
||||
)
|
||||
.unwrap();
|
||||
request_tx.send(request).unwrap();
|
||||
});
|
||||
let config = OtlpEnvSettings {
|
||||
traces_endpoint: Some(format!("http://{address}/v1/traces")),
|
||||
traces_protocol: Some("http/protobuf".to_owned()),
|
||||
traces_timeout: Some("2000".to_owned()),
|
||||
traces_headers: Some(String::new()),
|
||||
generic_headers: Some("authorization=Bearer%20canary-token".to_owned()),
|
||||
max_queue_size: Some("16".to_owned()),
|
||||
max_export_batch_size: Some("8".to_owned()),
|
||||
scheduled_delay: Some("10".to_owned()),
|
||||
batch_export_timeout: Some("2000".to_owned()),
|
||||
..OtlpEnvSettings::default()
|
||||
}
|
||||
.into_config()
|
||||
.unwrap();
|
||||
let identity = ServiceIdentity::try_new("admin-api", "0.3.1", "integration-test").unwrap();
|
||||
let (provider, tracer) = build_tracer_provider(&identity, &config).unwrap().unwrap();
|
||||
let mut span = tracer.start("http.request");
|
||||
let trace_id = span.span_context().trace_id().to_bytes();
|
||||
span.set_attribute(KeyValue::new("request_id", "req_otlp_contract"));
|
||||
span.set_attribute(KeyValue::new("authorization", "Bearer canary-span-secret"));
|
||||
span.add_event(
|
||||
"canary-span-event",
|
||||
vec![KeyValue::new("payload", "canary-span-secret")],
|
||||
);
|
||||
span.set_status(Status::error("canary-span-secret"));
|
||||
span.end();
|
||||
|
||||
provider.force_flush().unwrap();
|
||||
provider.shutdown().unwrap();
|
||||
let request = request_rx.recv_timeout(Duration::from_secs(2)).unwrap();
|
||||
server.join().unwrap();
|
||||
let (headers, body) = split_http_request(&request);
|
||||
|
||||
assert!(headers.contains("POST /v1/traces HTTP/1.1"));
|
||||
assert!(
|
||||
headers
|
||||
.to_ascii_lowercase()
|
||||
.contains("content-type: application/x-protobuf")
|
||||
);
|
||||
assert!(
|
||||
headers
|
||||
.to_ascii_lowercase()
|
||||
.contains("authorization: bearer canary-token")
|
||||
);
|
||||
let export = ExportTraceServiceRequest::decode(body).unwrap();
|
||||
let resource_spans = export.resource_spans.first().unwrap();
|
||||
let attributes = &resource_spans.resource.as_ref().unwrap().attributes;
|
||||
assert_eq!(
|
||||
string_attribute(attributes, "service.name"),
|
||||
Some("admin-api")
|
||||
);
|
||||
assert_eq!(
|
||||
string_attribute(attributes, "service.version"),
|
||||
Some("0.3.1")
|
||||
);
|
||||
assert_eq!(
|
||||
string_attribute(attributes, "deployment.environment.name"),
|
||||
Some("integration-test")
|
||||
);
|
||||
assert_eq!(
|
||||
resource_spans.scope_spans[0].spans[0].trace_id.as_slice(),
|
||||
trace_id
|
||||
);
|
||||
let exported_span = &resource_spans.scope_spans[0].spans[0];
|
||||
assert_eq!(exported_span.name, "http.request");
|
||||
assert_eq!(
|
||||
string_attribute(&exported_span.attributes, "request_id"),
|
||||
Some("req_otlp_contract")
|
||||
);
|
||||
assert!(
|
||||
exported_span
|
||||
.attributes
|
||||
.iter()
|
||||
.all(|attribute| attribute.key != "authorization")
|
||||
);
|
||||
assert!(exported_span.events.is_empty());
|
||||
assert_eq!(
|
||||
exported_span
|
||||
.status
|
||||
.as_ref()
|
||||
.map(|status| status.message.as_str()),
|
||||
Some("")
|
||||
);
|
||||
assert!(
|
||||
!body
|
||||
.windows(b"canary-span-secret".len())
|
||||
.any(|window| { window == b"canary-span-secret" })
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "current_thread")]
|
||||
async fn unavailable_receiver_does_not_change_product_result() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let server = thread::spawn(move || {
|
||||
let (stream, _) = listener.accept().unwrap();
|
||||
drop(stream);
|
||||
});
|
||||
let config = OtlpTraceConfig::try_new(
|
||||
Some(format!("http://{address}/v1/traces")),
|
||||
Some("http/protobuf".to_owned()),
|
||||
Duration::from_millis(250),
|
||||
OtlpBatchConfig::try_new(8, 4, Duration::from_millis(10), Duration::from_millis(250))
|
||||
.unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
let identity = ServiceIdentity::try_new("mcp-server", "0.3.1", "fault-test").unwrap();
|
||||
let (provider, tracer) = build_tracer_provider(&identity, &config).unwrap().unwrap();
|
||||
let subscriber =
|
||||
tracing_subscriber::registry().with(tracing_opentelemetry::layer().with_tracer(tracer));
|
||||
let dispatch = tracing::Dispatch::new(subscriber);
|
||||
let _dispatch_guard = tracing::dispatcher::set_default(&dispatch);
|
||||
let app = Router::new()
|
||||
.route("/product", get(|| async { "product-success" }))
|
||||
.layer(axum::middleware::from_fn(trace_product_request));
|
||||
|
||||
let response = app
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/product")
|
||||
.body(Body::empty())
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let status = response.status();
|
||||
let body = to_bytes(response.into_body(), 64).await.unwrap();
|
||||
|
||||
assert_eq!(status, axum::http::StatusCode::OK);
|
||||
assert_eq!(body.as_ref(), b"product-success");
|
||||
assert!(provider.force_flush().is_err());
|
||||
let _ = provider.shutdown();
|
||||
server.join().unwrap();
|
||||
}
|
||||
|
||||
async fn trace_product_request(request: Request, next: Next) -> Response {
|
||||
next.run(request)
|
||||
.instrument(info_span!(target: "crank::trace", "http.request"))
|
||||
.await
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn hanging_receiver_respects_the_stricter_export_timeout() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let server = thread::spawn(move || {
|
||||
let (stream, _) = listener.accept().unwrap();
|
||||
thread::sleep(Duration::from_millis(750));
|
||||
drop(stream);
|
||||
});
|
||||
let config = OtlpTraceConfig::try_new(
|
||||
Some(format!("http://{address}/v1/traces")),
|
||||
Some("http/protobuf".to_owned()),
|
||||
Duration::from_secs(2),
|
||||
OtlpBatchConfig::try_new(8, 4, Duration::from_millis(10), Duration::from_millis(100))
|
||||
.unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
let identity = ServiceIdentity::try_new("admin-api", "0.3.1", "timeout-test").unwrap();
|
||||
let (provider, tracer) = build_tracer_provider(&identity, &config).unwrap().unwrap();
|
||||
let mut span = tracer.start("http.request");
|
||||
span.end();
|
||||
let started_at = Instant::now();
|
||||
|
||||
assert!(provider.force_flush().is_err());
|
||||
assert!(started_at.elapsed() < Duration::from_millis(500));
|
||||
let _ = provider.shutdown();
|
||||
server.join().unwrap();
|
||||
}
|
||||
|
||||
fn read_http_request(stream: &mut std::net::TcpStream) -> Vec<u8> {
|
||||
stream
|
||||
.set_read_timeout(Some(Duration::from_secs(2)))
|
||||
.unwrap();
|
||||
let mut request = Vec::new();
|
||||
let mut buffer = [0_u8; 4096];
|
||||
loop {
|
||||
let read = stream.read(&mut buffer).unwrap();
|
||||
request.extend_from_slice(&buffer[..read]);
|
||||
let Some(header_end) = find_bytes(&request, b"\r\n\r\n") else {
|
||||
continue;
|
||||
};
|
||||
let headers = String::from_utf8_lossy(&request[..header_end]);
|
||||
let content_length = headers
|
||||
.lines()
|
||||
.find_map(|line| {
|
||||
let (name, value) = line.split_once(':')?;
|
||||
name.eq_ignore_ascii_case("content-length")
|
||||
.then(|| value.trim().parse::<usize>().ok())
|
||||
.flatten()
|
||||
})
|
||||
.unwrap_or(0);
|
||||
if request.len() >= header_end + 4 + content_length {
|
||||
return request;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn split_http_request(request: &[u8]) -> (&str, &[u8]) {
|
||||
let header_end = find_bytes(request, b"\r\n\r\n").unwrap();
|
||||
(
|
||||
std::str::from_utf8(&request[..header_end]).unwrap(),
|
||||
&request[header_end + 4..],
|
||||
)
|
||||
}
|
||||
|
||||
fn string_attribute<'a>(
|
||||
attributes: &'a [opentelemetry_proto::tonic::common::v1::KeyValue],
|
||||
key: &str,
|
||||
) -> Option<&'a str> {
|
||||
attributes.iter().find_map(|attribute| {
|
||||
let value = attribute.value.as_ref()?.value.as_ref()?;
|
||||
(attribute.key == key)
|
||||
.then_some(value)
|
||||
.and_then(|value| match value {
|
||||
any_value::Value::StringValue(value) => Some(value.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn find_bytes(haystack: &[u8], needle: &[u8]) -> Option<usize> {
|
||||
haystack
|
||||
.windows(needle.len())
|
||||
.position(|candidate| candidate == needle)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,283 @@
|
||||
use std::{env, net::SocketAddr};
|
||||
|
||||
use axum::{
|
||||
Router,
|
||||
extract::{Request, State},
|
||||
http::{
|
||||
HeaderMap, StatusCode,
|
||||
header::{self, HeaderValue},
|
||||
},
|
||||
middleware::{self, Next},
|
||||
response::{IntoResponse, Response},
|
||||
routing::get,
|
||||
};
|
||||
use metrics_exporter_prometheus::{PrometheusBuilder, PrometheusHandle, PrometheusRecorder};
|
||||
use sha2::{Digest, Sha256};
|
||||
use subtle::ConstantTimeEq;
|
||||
use thiserror::Error;
|
||||
use tokio::net::TcpListener;
|
||||
|
||||
use crate::{DURATION_BUCKETS_SECONDS, ServiceIdentity};
|
||||
|
||||
const METRICS_ENABLED_ENV: &str = "CRANK_METRICS_ENABLED";
|
||||
const METRICS_TOKEN_ENV: &str = "CRANK_METRICS_BEARER_TOKEN";
|
||||
const PROMETHEUS_CONTENT_TYPE: &str = "text/plain; version=0.0.4; charset=utf-8";
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct MetricsConfig {
|
||||
enabled: bool,
|
||||
bind_addr: SocketAddr,
|
||||
token_digest: Option<[u8; 32]>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for MetricsConfig {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("MetricsConfig")
|
||||
.field("enabled", &self.enabled)
|
||||
.field("bind_addr", &self.bind_addr)
|
||||
.field("authentication_configured", &self.token_digest.is_some())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl MetricsConfig {
|
||||
pub fn new(
|
||||
enabled: bool,
|
||||
bind_addr: SocketAddr,
|
||||
bearer_token: Option<String>,
|
||||
) -> Result<Self, MetricsConfigError> {
|
||||
let token_digest = bearer_token
|
||||
.filter(|token| !token.is_empty())
|
||||
.map(|token| token_digest(token.as_bytes()));
|
||||
|
||||
if enabled && !bind_addr.ip().is_loopback() && token_digest.is_none() {
|
||||
return Err(MetricsConfigError::MissingTokenForExternalBind);
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
enabled,
|
||||
bind_addr,
|
||||
token_digest,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn from_env(
|
||||
bind_env: &'static str,
|
||||
default_bind: SocketAddr,
|
||||
) -> Result<Self, MetricsConfigError> {
|
||||
let enabled = parse_enabled(env::var(METRICS_ENABLED_ENV))?;
|
||||
let bind_addr = match env::var(bind_env) {
|
||||
Ok(raw) => raw
|
||||
.parse()
|
||||
.map_err(|_| MetricsConfigError::InvalidBindAddress { field: bind_env })?,
|
||||
Err(env::VarError::NotPresent) => default_bind,
|
||||
Err(env::VarError::NotUnicode(_)) => {
|
||||
return Err(MetricsConfigError::InvalidEnvironmentEncoding { field: bind_env });
|
||||
}
|
||||
};
|
||||
let bearer_token = match env::var(METRICS_TOKEN_ENV) {
|
||||
Ok(token) => Some(token),
|
||||
Err(env::VarError::NotPresent) => None,
|
||||
Err(env::VarError::NotUnicode(_)) => {
|
||||
return Err(MetricsConfigError::InvalidEnvironmentEncoding {
|
||||
field: METRICS_TOKEN_ENV,
|
||||
});
|
||||
}
|
||||
};
|
||||
|
||||
Self::new(enabled, bind_addr, bearer_token)
|
||||
}
|
||||
|
||||
pub fn enabled(&self) -> bool {
|
||||
self.enabled
|
||||
}
|
||||
|
||||
pub fn bind_addr(&self) -> SocketAddr {
|
||||
self.bind_addr
|
||||
}
|
||||
|
||||
pub fn requires_authentication(&self) -> bool {
|
||||
!self.bind_addr.ip().is_loopback()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum MetricsConfigError {
|
||||
#[error("metrics environment variable is not valid UTF-8: {field}")]
|
||||
InvalidEnvironmentEncoding { field: &'static str },
|
||||
#[error("metrics bind address is invalid: {field}")]
|
||||
InvalidBindAddress { field: &'static str },
|
||||
#[error("metrics enabled flag must be one of true, false, 1, 0")]
|
||||
InvalidEnabledFlag,
|
||||
#[error("external metrics bind requires a bearer token")]
|
||||
MissingTokenForExternalBind,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct MetricsState {
|
||||
handle: PrometheusHandle,
|
||||
token_digest: Option<[u8; 32]>,
|
||||
requires_authentication: bool,
|
||||
}
|
||||
|
||||
pub struct MetricsSurface {
|
||||
config: MetricsConfig,
|
||||
state: MetricsState,
|
||||
_recorder: Option<PrometheusRecorder>,
|
||||
}
|
||||
|
||||
impl MetricsSurface {
|
||||
pub(crate) fn new(config: MetricsConfig, handle: PrometheusHandle) -> Self {
|
||||
Self {
|
||||
state: MetricsState {
|
||||
handle,
|
||||
token_digest: config.token_digest,
|
||||
requires_authentication: config.requires_authentication(),
|
||||
},
|
||||
config,
|
||||
_recorder: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn for_test(
|
||||
config: MetricsConfig,
|
||||
identity: ServiceIdentity,
|
||||
) -> Result<Self, MetricsSurfaceError> {
|
||||
let recorder = prometheus_builder(&identity)?.build_recorder();
|
||||
let handle = recorder.handle();
|
||||
let mut surface = Self::new(config, handle);
|
||||
surface._recorder = Some(recorder);
|
||||
Ok(surface)
|
||||
}
|
||||
|
||||
pub fn router(&self) -> Router {
|
||||
Router::new()
|
||||
.route("/metrics", get(render_metrics))
|
||||
.route("/health", get(metrics_health))
|
||||
.layer(middleware::from_fn_with_state(
|
||||
self.state.clone(),
|
||||
authorize_metrics,
|
||||
))
|
||||
.with_state(self.state.clone())
|
||||
}
|
||||
|
||||
pub async fn bind(self) -> Result<MetricsServer, MetricsServeError> {
|
||||
let listener = TcpListener::bind(self.config.bind_addr)
|
||||
.await
|
||||
.map_err(|_| MetricsServeError::Bind)?;
|
||||
Ok(MetricsServer {
|
||||
listener,
|
||||
router: self.router(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub struct MetricsServer {
|
||||
listener: TcpListener,
|
||||
router: Router,
|
||||
}
|
||||
|
||||
impl MetricsServer {
|
||||
pub async fn serve(self) -> Result<(), MetricsServeError> {
|
||||
axum::serve(self.listener, self.router)
|
||||
.await
|
||||
.map_err(|_| MetricsServeError::Serve)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum MetricsSurfaceError {
|
||||
#[error("failed to configure Prometheus recorder")]
|
||||
RecorderConfiguration,
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum MetricsServeError {
|
||||
#[error("failed to bind metrics listener")]
|
||||
Bind,
|
||||
#[error("metrics listener stopped unexpectedly")]
|
||||
Serve,
|
||||
}
|
||||
|
||||
pub(crate) fn install_prometheus_recorder(
|
||||
identity: &ServiceIdentity,
|
||||
) -> Result<PrometheusHandle, MetricsSurfaceError> {
|
||||
prometheus_builder(identity)?
|
||||
.install_recorder()
|
||||
.map_err(|_| MetricsSurfaceError::RecorderConfiguration)
|
||||
}
|
||||
|
||||
fn prometheus_builder(
|
||||
identity: &ServiceIdentity,
|
||||
) -> Result<PrometheusBuilder, MetricsSurfaceError> {
|
||||
PrometheusBuilder::new()
|
||||
.set_buckets(DURATION_BUCKETS_SECONDS)
|
||||
.map(|builder| {
|
||||
builder
|
||||
.add_global_label("service", identity.service())
|
||||
.add_global_label("version", identity.version())
|
||||
.add_global_label("environment", identity.environment())
|
||||
})
|
||||
.map_err(|_| MetricsSurfaceError::RecorderConfiguration)
|
||||
}
|
||||
|
||||
async fn render_metrics(State(state): State<MetricsState>) -> Response {
|
||||
let mut response = state.handle.render().into_response();
|
||||
response.headers_mut().insert(
|
||||
header::CONTENT_TYPE,
|
||||
HeaderValue::from_static(PROMETHEUS_CONTENT_TYPE),
|
||||
);
|
||||
response
|
||||
}
|
||||
|
||||
async fn metrics_health() -> impl IntoResponse {
|
||||
(StatusCode::OK, "ok\n")
|
||||
}
|
||||
|
||||
async fn authorize_metrics(
|
||||
State(state): State<MetricsState>,
|
||||
request: Request,
|
||||
next: Next,
|
||||
) -> Response {
|
||||
if !state.requires_authentication {
|
||||
return next.run(request).await;
|
||||
}
|
||||
|
||||
let authorized = bearer_token(request.headers())
|
||||
.map(token_digest)
|
||||
.zip(state.token_digest)
|
||||
.is_some_and(|(actual, expected)| bool::from(actual.ct_eq(&expected)));
|
||||
|
||||
if authorized {
|
||||
next.run(request).await
|
||||
} else {
|
||||
StatusCode::UNAUTHORIZED.into_response()
|
||||
}
|
||||
}
|
||||
|
||||
fn bearer_token(headers: &HeaderMap) -> Option<&[u8]> {
|
||||
headers
|
||||
.get(header::AUTHORIZATION)?
|
||||
.as_bytes()
|
||||
.strip_prefix(b"Bearer ")
|
||||
.filter(|token| !token.is_empty())
|
||||
}
|
||||
|
||||
fn token_digest(token: &[u8]) -> [u8; 32] {
|
||||
Sha256::digest(token).into()
|
||||
}
|
||||
|
||||
fn parse_enabled(value: Result<String, env::VarError>) -> Result<bool, MetricsConfigError> {
|
||||
match value {
|
||||
Ok(raw) => match raw.to_ascii_lowercase().as_str() {
|
||||
"true" | "1" => Ok(true),
|
||||
"false" | "0" => Ok(false),
|
||||
_ => Err(MetricsConfigError::InvalidEnabledFlag),
|
||||
},
|
||||
Err(env::VarError::NotPresent) => Ok(true),
|
||||
Err(env::VarError::NotUnicode(_)) => Err(MetricsConfigError::InvalidEnvironmentEncoding {
|
||||
field: METRICS_ENABLED_ENV,
|
||||
}),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,70 @@
|
||||
use axum::http::{HeaderMap, HeaderName, HeaderValue};
|
||||
use opentelemetry::{
|
||||
Context, global,
|
||||
propagation::{Extractor, Injector},
|
||||
trace::TraceContextExt,
|
||||
};
|
||||
use opentelemetry_sdk::propagation::TraceContextPropagator;
|
||||
use tracing::Span;
|
||||
use tracing_opentelemetry::OpenTelemetrySpanExt;
|
||||
|
||||
pub fn set_remote_trace_parent(span: &Span, headers: &HeaderMap) -> bool {
|
||||
let context =
|
||||
global::get_text_map_propagator(|propagator| propagator.extract(&HeaderExtractor(headers)));
|
||||
let span_context = context.span().span_context().clone();
|
||||
if !span_context.is_valid() || !span_context.is_remote() {
|
||||
return false;
|
||||
}
|
||||
|
||||
span.set_parent(context).is_ok()
|
||||
}
|
||||
|
||||
pub fn inject_current_trace_context(headers: &mut HeaderMap) -> bool {
|
||||
let context = Span::current().context();
|
||||
if !context.span().span_context().is_valid() {
|
||||
return false;
|
||||
}
|
||||
|
||||
global::get_text_map_propagator(|propagator| {
|
||||
propagator.inject_context(&context, &mut HeaderInjector(headers));
|
||||
});
|
||||
true
|
||||
}
|
||||
|
||||
pub(crate) fn install_trace_context_propagator() {
|
||||
global::set_text_map_propagator(TraceContextPropagator::new());
|
||||
}
|
||||
|
||||
struct HeaderExtractor<'a>(&'a HeaderMap);
|
||||
|
||||
impl Extractor for HeaderExtractor<'_> {
|
||||
fn get(&self, key: &str) -> Option<&str> {
|
||||
self.0.get(key).and_then(|value| value.to_str().ok())
|
||||
}
|
||||
|
||||
fn keys(&self) -> Vec<&str> {
|
||||
self.0.keys().map(HeaderName::as_str).collect()
|
||||
}
|
||||
}
|
||||
|
||||
struct HeaderInjector<'a>(&'a mut HeaderMap);
|
||||
|
||||
impl Injector for HeaderInjector<'_> {
|
||||
fn set(&mut self, key: &str, value: String) {
|
||||
let Ok(name) = HeaderName::try_from(key) else {
|
||||
return;
|
||||
};
|
||||
let Ok(value) = HeaderValue::try_from(value) else {
|
||||
return;
|
||||
};
|
||||
self.0.insert(name, value);
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn current_trace_id() -> Option<String> {
|
||||
let context: Context = Span::current().context();
|
||||
let span_context = context.span().span_context().clone();
|
||||
span_context
|
||||
.is_valid()
|
||||
.then(|| span_context.trace_id().to_string())
|
||||
}
|
||||
@@ -0,0 +1,292 @@
|
||||
use serde_json::{Map, Value};
|
||||
use thiserror::Error;
|
||||
|
||||
pub const REDACTED_MARKER: &str = "[REDACTED]";
|
||||
const TRUNCATED_MARKER: &str = "[TRUNCATED]";
|
||||
const MIN_EVENT_BYTES: usize = 512;
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub struct RedactionLimits {
|
||||
pub max_string_bytes: usize,
|
||||
pub max_array_items: usize,
|
||||
pub max_object_fields: usize,
|
||||
pub max_depth: usize,
|
||||
pub max_event_bytes: usize,
|
||||
}
|
||||
|
||||
impl Default for RedactionLimits {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
max_string_bytes: 1024,
|
||||
max_array_items: 32,
|
||||
max_object_fields: 64,
|
||||
max_depth: 8,
|
||||
max_event_bytes: 16 * 1024,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl RedactionLimits {
|
||||
pub fn validate(self) -> Result<(), RedactionLimitsError> {
|
||||
for (field, value, minimum) in [
|
||||
(
|
||||
"max_string_bytes",
|
||||
self.max_string_bytes,
|
||||
TRUNCATED_MARKER.len(),
|
||||
),
|
||||
("max_array_items", self.max_array_items, 1),
|
||||
("max_object_fields", self.max_object_fields, 1),
|
||||
("max_depth", self.max_depth, 1),
|
||||
("max_event_bytes", self.max_event_bytes, MIN_EVENT_BYTES),
|
||||
] {
|
||||
if value < minimum {
|
||||
return Err(RedactionLimitsError::TooSmall { field, minimum });
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum RedactionLimitsError {
|
||||
#[error("invalid redaction limit {field}: minimum is {minimum}")]
|
||||
TooSmall { field: &'static str, minimum: usize },
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum SafeJsonError {
|
||||
#[error(transparent)]
|
||||
InvalidLimits(#[from] RedactionLimitsError),
|
||||
#[error(transparent)]
|
||||
Serialization(#[from] serde_json::Error),
|
||||
}
|
||||
|
||||
pub fn redact_value(value: &Value, limits: RedactionLimits) -> Value {
|
||||
redact_at_depth(value, limits, 0)
|
||||
}
|
||||
|
||||
pub fn safe_json(value: &Value, limits: RedactionLimits) -> Result<String, SafeJsonError> {
|
||||
limits.validate()?;
|
||||
let serialized = serde_json::to_string(&redact_value(value, limits))?;
|
||||
if serialized.len() <= limits.max_event_bytes {
|
||||
return Ok(serialized);
|
||||
}
|
||||
|
||||
Ok(serde_json::to_string(&serde_json::json!({
|
||||
"truncated": true
|
||||
}))?)
|
||||
}
|
||||
|
||||
pub(crate) fn truncate_string(value: &str, max_bytes: usize) -> String {
|
||||
if value.len() <= max_bytes {
|
||||
return value.to_owned();
|
||||
}
|
||||
if max_bytes == 0 {
|
||||
return String::new();
|
||||
}
|
||||
|
||||
let marker = if max_bytes >= TRUNCATED_MARKER.len() {
|
||||
TRUNCATED_MARKER
|
||||
} else {
|
||||
""
|
||||
};
|
||||
let content_budget = max_bytes.saturating_sub(marker.len());
|
||||
let mut boundary = content_budget.min(value.len());
|
||||
while boundary > 0 && !value.is_char_boundary(boundary) {
|
||||
boundary -= 1;
|
||||
}
|
||||
|
||||
let mut truncated = String::with_capacity(max_bytes);
|
||||
truncated.push_str(&value[..boundary]);
|
||||
if marker.is_empty() {
|
||||
let mut marker_boundary = max_bytes.min(TRUNCATED_MARKER.len());
|
||||
while marker_boundary > 0 && !TRUNCATED_MARKER.is_char_boundary(marker_boundary) {
|
||||
marker_boundary -= 1;
|
||||
}
|
||||
truncated.clear();
|
||||
truncated.push_str(&TRUNCATED_MARKER[..marker_boundary]);
|
||||
} else {
|
||||
truncated.push_str(marker);
|
||||
}
|
||||
truncated
|
||||
}
|
||||
|
||||
fn redact_at_depth(value: &Value, limits: RedactionLimits, depth: usize) -> Value {
|
||||
if depth >= limits.max_depth {
|
||||
return Value::String(TRUNCATED_MARKER.to_owned());
|
||||
}
|
||||
|
||||
match value {
|
||||
Value::Null | Value::Bool(_) | Value::Number(_) => value.clone(),
|
||||
Value::String(value) => Value::String(truncate_string(value, limits.max_string_bytes)),
|
||||
Value::Array(values) => redact_array(values, limits, depth),
|
||||
Value::Object(values) => redact_object(values, limits, depth),
|
||||
}
|
||||
}
|
||||
|
||||
fn redact_array(values: &[Value], limits: RedactionLimits, depth: usize) -> Value {
|
||||
if limits.max_array_items == 0 {
|
||||
return Value::Array(Vec::new());
|
||||
}
|
||||
|
||||
let truncated = values.len() > limits.max_array_items;
|
||||
let value_limit = if truncated {
|
||||
limits.max_array_items.saturating_sub(1)
|
||||
} else {
|
||||
limits.max_array_items
|
||||
};
|
||||
let mut output: Vec<_> = values
|
||||
.iter()
|
||||
.take(value_limit)
|
||||
.map(|value| redact_at_depth(value, limits, depth + 1))
|
||||
.collect();
|
||||
if truncated {
|
||||
output.push(Value::String(TRUNCATED_MARKER.to_owned()));
|
||||
}
|
||||
Value::Array(output)
|
||||
}
|
||||
|
||||
fn redact_object(values: &Map<String, Value>, limits: RedactionLimits, depth: usize) -> Value {
|
||||
if limits.max_object_fields == 0 {
|
||||
return Value::Object(Map::new());
|
||||
}
|
||||
|
||||
let truncated = values.len() > limits.max_object_fields;
|
||||
let value_limit = if truncated {
|
||||
limits.max_object_fields.saturating_sub(1)
|
||||
} else {
|
||||
limits.max_object_fields
|
||||
};
|
||||
let mut output = Map::new();
|
||||
|
||||
for (index, (key, value)) in values.iter().take(value_limit).enumerate() {
|
||||
let cleaned = if is_sensitive_key(key) {
|
||||
Value::String(REDACTED_MARKER.to_owned())
|
||||
} else if is_url_key(key) {
|
||||
value
|
||||
.as_str()
|
||||
.map(sanitize_url)
|
||||
.map(|value| truncate_string(&value, limits.max_string_bytes))
|
||||
.map(Value::String)
|
||||
.unwrap_or_else(|| redact_at_depth(value, limits, depth + 1))
|
||||
} else {
|
||||
redact_at_depth(value, limits, depth + 1)
|
||||
};
|
||||
output.insert(
|
||||
bounded_object_key(key, limits.max_string_bytes, index),
|
||||
cleaned,
|
||||
);
|
||||
}
|
||||
if truncated {
|
||||
output.insert(
|
||||
"_truncated".to_owned(),
|
||||
Value::String(TRUNCATED_MARKER.to_owned()),
|
||||
);
|
||||
}
|
||||
|
||||
Value::Object(output)
|
||||
}
|
||||
|
||||
fn normalized_key(key: &str) -> String {
|
||||
key.bytes()
|
||||
.filter(|byte| !matches!(byte, b'_' | b'-' | b'.'))
|
||||
.map(|byte| byte.to_ascii_lowercase() as char)
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn is_sensitive_key(key: &str) -> bool {
|
||||
let key = normalized_key(key);
|
||||
let exact_match = matches!(
|
||||
key.as_str(),
|
||||
"password"
|
||||
| "passwd"
|
||||
| "secret"
|
||||
| "token"
|
||||
| "apikey"
|
||||
| "accesskey"
|
||||
| "secretkey"
|
||||
| "authorization"
|
||||
| "proxyauthorization"
|
||||
| "cookie"
|
||||
| "setcookie"
|
||||
| "query"
|
||||
| "querystring"
|
||||
| "rawquery"
|
||||
| "urlquery"
|
||||
| "payload"
|
||||
| "body"
|
||||
| "requestbody"
|
||||
| "arguments"
|
||||
| "result"
|
||||
| "response"
|
||||
| "context"
|
||||
| "error"
|
||||
| "errormessage"
|
||||
);
|
||||
let contains_high_risk_name = [
|
||||
"password",
|
||||
"passwd",
|
||||
"secret",
|
||||
"token",
|
||||
"apikey",
|
||||
"accesskey",
|
||||
"authorization",
|
||||
"cookie",
|
||||
]
|
||||
.iter()
|
||||
.any(|part| key.contains(part));
|
||||
|
||||
exact_match
|
||||
|| contains_high_risk_name
|
||||
|| key.ends_with("payload")
|
||||
|| key.ends_with("body")
|
||||
|| key.ends_with("arguments")
|
||||
|| key.ends_with("result")
|
||||
|| key.ends_with("response")
|
||||
|| key.ends_with("query")
|
||||
|| key.ends_with("context")
|
||||
|| key.starts_with("query")
|
||||
}
|
||||
|
||||
fn is_url_key(key: &str) -> bool {
|
||||
let key = normalized_key(key);
|
||||
matches!(
|
||||
key.as_str(),
|
||||
"url" | "uri" | "endpoint" | "endpointurl" | "endpointuri" | "requesturl" | "targeturl"
|
||||
) || key.ends_with("url")
|
||||
|| key.ends_with("uri")
|
||||
|| key.ends_with("endpoint")
|
||||
}
|
||||
|
||||
fn sanitize_url(value: &str) -> String {
|
||||
let without_query = value
|
||||
.find(['?', '#'])
|
||||
.map(|index| &value[..index])
|
||||
.unwrap_or(value);
|
||||
let Some(scheme_end) = without_query.find("://") else {
|
||||
return without_query.to_owned();
|
||||
};
|
||||
let authority_start = scheme_end + 3;
|
||||
let authority_end = without_query[authority_start..]
|
||||
.find('/')
|
||||
.map(|index| authority_start + index)
|
||||
.unwrap_or(without_query.len());
|
||||
let authority = &without_query[authority_start..authority_end];
|
||||
let Some(userinfo_end) = authority.rfind('@') else {
|
||||
return without_query.to_owned();
|
||||
};
|
||||
|
||||
format!(
|
||||
"{}{}",
|
||||
&without_query[..authority_start],
|
||||
&without_query[authority_start + userinfo_end + 1..]
|
||||
)
|
||||
}
|
||||
|
||||
fn bounded_object_key(key: &str, max_bytes: usize, index: usize) -> String {
|
||||
if key.len() <= max_bytes {
|
||||
return key.to_owned();
|
||||
}
|
||||
|
||||
truncate_string(&format!("_truncated_key_{index}"), max_bytes)
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
use serde::Serialize;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(crate) struct LogEnvelope {
|
||||
pub timestamp: String,
|
||||
pub level: String,
|
||||
pub service: String,
|
||||
pub version: String,
|
||||
pub environment: String,
|
||||
pub target: String,
|
||||
pub event: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub request_id: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub trace_id: Option<String>,
|
||||
pub fields: Map<String, Value>,
|
||||
}
|
||||
@@ -0,0 +1,42 @@
|
||||
use crank_observability::RequestId;
|
||||
use uuid::Version;
|
||||
|
||||
#[test]
|
||||
fn preserves_valid_opaque_request_id() {
|
||||
let request_id = RequestId::resolve(Some("req_test-123/abc"));
|
||||
|
||||
assert_eq!(request_id.as_str(), "req_test-123/abc");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn replaces_missing_and_invalid_values_with_uuid_v7() {
|
||||
for candidate in [
|
||||
None,
|
||||
Some(""),
|
||||
Some("bad value"),
|
||||
Some(" leading"),
|
||||
Some("trailing "),
|
||||
Some("bad,value"),
|
||||
Some("bad;value"),
|
||||
Some("я"),
|
||||
] {
|
||||
let request_id = RequestId::resolve(candidate);
|
||||
let parsed = uuid::Uuid::parse_str(request_id.as_str()).expect("generated UUID");
|
||||
|
||||
assert_eq!(parsed.get_version(), Some(Version::SortRand));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_values_over_the_shared_limit() {
|
||||
let oversized = "x".repeat(RequestId::MAX_LEN + 1);
|
||||
let request_id = RequestId::resolve(Some(&oversized));
|
||||
|
||||
assert_ne!(request_id.as_str(), oversized);
|
||||
assert_eq!(
|
||||
uuid::Uuid::parse_str(request_id.as_str())
|
||||
.expect("generated UUID")
|
||||
.get_version(),
|
||||
Some(Version::SortRand)
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
use crank_observability::{CriticalErrorCategory, SentryConfig, SentryConfigError};
|
||||
|
||||
#[test]
|
||||
fn missing_or_blank_dsn_disables_critical_error_channel() {
|
||||
assert!(!SentryConfig::parse(None).expect("missing DSN").enabled());
|
||||
assert!(!SentryConfig::parse(Some("")).expect("empty DSN").enabled());
|
||||
assert!(
|
||||
!SentryConfig::parse(Some(" "))
|
||||
.expect("blank DSN")
|
||||
.enabled()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_explicit_dsn_is_rejected_without_echoing_the_value() {
|
||||
let secret_value = "not-a-dsn?token=control-secret";
|
||||
let error = SentryConfig::parse(Some(secret_value)).expect_err("invalid DSN must fail");
|
||||
|
||||
assert!(matches!(error, SentryConfigError::InvalidDsn));
|
||||
assert!(!error.to_string().contains(secret_value));
|
||||
assert!(!error.to_string().contains("control-secret"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn critical_error_categories_are_closed_and_stable() {
|
||||
assert_eq!(CriticalErrorCategory::Panic.as_str(), "panic");
|
||||
assert_eq!(CriticalErrorCategory::Startup.as_str(), "startup");
|
||||
assert_eq!(CriticalErrorCategory::Internal.as_str(), "internal");
|
||||
assert_eq!(
|
||||
CriticalErrorCategory::DataIntegrity.as_str(),
|
||||
"data_integrity"
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,92 @@
|
||||
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
|
||||
|
||||
use axum::{
|
||||
Router,
|
||||
body::{Body, to_bytes},
|
||||
http::{Request, StatusCode},
|
||||
middleware,
|
||||
routing::get,
|
||||
};
|
||||
use crank_observability::{
|
||||
MetricsConfig, ObservabilityConfig, RedactionLimits, ServiceIdentity, record_http_request,
|
||||
};
|
||||
use tower::ServiceExt;
|
||||
|
||||
#[tokio::test]
|
||||
async fn http_metrics_use_matched_routes_and_closed_labels() {
|
||||
let identity =
|
||||
ServiceIdentity::try_new("metrics-test", "0.3.1", "test").expect("valid identity");
|
||||
let lifecycle = crank_observability::init(ObservabilityConfig::new(
|
||||
identity,
|
||||
"off",
|
||||
RedactionLimits::default(),
|
||||
))
|
||||
.expect("observability lifecycle");
|
||||
let config = MetricsConfig::new(
|
||||
true,
|
||||
SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 9464),
|
||||
None,
|
||||
)
|
||||
.expect("loopback metrics");
|
||||
let metrics = lifecycle.metrics_surface(config).router();
|
||||
let app = Router::new()
|
||||
.route(
|
||||
"/documents/{document_id}",
|
||||
get(|| async { StatusCode::NO_CONTENT }),
|
||||
)
|
||||
.layer(middleware::from_fn(record_http_request));
|
||||
|
||||
let sensitive_path_segment = "customer-secret-document-id";
|
||||
for index in 0..100 {
|
||||
let response = app
|
||||
.clone()
|
||||
.oneshot(
|
||||
Request::get(format!("/documents/{sensitive_path_segment}-{index}"))
|
||||
.body(Body::empty())
|
||||
.expect("request"),
|
||||
)
|
||||
.await
|
||||
.expect("response");
|
||||
assert_eq!(response.status(), StatusCode::NO_CONTENT);
|
||||
}
|
||||
|
||||
let response = app
|
||||
.oneshot(
|
||||
Request::get("/unknown/customer-controlled-path")
|
||||
.body(Body::empty())
|
||||
.expect("request"),
|
||||
)
|
||||
.await
|
||||
.expect("response");
|
||||
assert_eq!(response.status(), StatusCode::NOT_FOUND);
|
||||
|
||||
let response = metrics
|
||||
.oneshot(
|
||||
Request::get("/metrics")
|
||||
.body(Body::empty())
|
||||
.expect("request"),
|
||||
)
|
||||
.await
|
||||
.expect("metrics response");
|
||||
let body = to_bytes(response.into_body(), 1024 * 1024)
|
||||
.await
|
||||
.expect("bounded metrics body");
|
||||
let body = String::from_utf8(body.to_vec()).expect("utf-8 metrics");
|
||||
|
||||
assert!(body.contains("crank_http_requests_total"));
|
||||
assert!(body.contains("route=\"/documents/{document_id}\""));
|
||||
assert!(body.contains("method=\"GET\""));
|
||||
assert!(body.contains("status_class=\"2xx\""));
|
||||
assert!(body.contains("crank_http_request_duration_seconds_bucket"));
|
||||
assert!(!body.contains(sensitive_path_segment));
|
||||
assert_eq!(
|
||||
body.lines()
|
||||
.filter(|line| {
|
||||
line.starts_with("crank_http_requests_total{")
|
||||
&& line.contains("route=\"/documents/{document_id}\"")
|
||||
})
|
||||
.count(),
|
||||
1,
|
||||
"different entity ids must not create additional series"
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
use crank_observability::{
|
||||
OperationalIncident, operational_incident_total, record_operational_incident,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn history_loss_counter_has_no_dynamic_dimensions() {
|
||||
let before = operational_incident_total(OperationalIncident::InvocationHistoryLost);
|
||||
|
||||
record_operational_incident(OperationalIncident::InvocationHistoryLost);
|
||||
|
||||
assert!(
|
||||
operational_incident_total(OperationalIncident::InvocationHistoryLost) > before,
|
||||
"the closed incident counter must increase"
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,392 @@
|
||||
use std::{
|
||||
io,
|
||||
sync::{Arc, Mutex},
|
||||
};
|
||||
|
||||
use crank_observability::{
|
||||
ObservabilityConfig, RedactionLimits, ServiceIdentity, build_subscriber, safe_json,
|
||||
};
|
||||
use serde_json::Value;
|
||||
use time::{OffsetDateTime, format_description::well_known::Rfc3339};
|
||||
use tracing_subscriber::fmt::MakeWriter;
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
struct SharedWriter {
|
||||
buffer: Arc<Mutex<Vec<u8>>>,
|
||||
}
|
||||
|
||||
impl SharedWriter {
|
||||
fn output(&self) -> String {
|
||||
String::from_utf8(self.buffer.lock().expect("test writer lock").clone())
|
||||
.expect("log output must be UTF-8")
|
||||
}
|
||||
}
|
||||
|
||||
impl<'a> MakeWriter<'a> for SharedWriter {
|
||||
type Writer = SharedWriterGuard;
|
||||
|
||||
fn make_writer(&'a self) -> Self::Writer {
|
||||
SharedWriterGuard {
|
||||
buffer: Arc::clone(&self.buffer),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct SharedWriterGuard {
|
||||
buffer: Arc<Mutex<Vec<u8>>>,
|
||||
}
|
||||
|
||||
impl io::Write for SharedWriterGuard {
|
||||
fn write(&mut self, bytes: &[u8]) -> io::Result<usize> {
|
||||
self.buffer
|
||||
.lock()
|
||||
.map_err(|_| io::Error::other("test writer lock poisoned"))?
|
||||
.extend_from_slice(bytes);
|
||||
Ok(bytes.len())
|
||||
}
|
||||
|
||||
fn flush(&mut self) -> io::Result<()> {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
fn capture(service: &'static str, emit: impl FnOnce()) -> Vec<Value> {
|
||||
capture_with_limits(service, RedactionLimits::default(), emit)
|
||||
}
|
||||
|
||||
fn capture_with_limits(
|
||||
service: &'static str,
|
||||
limits: RedactionLimits,
|
||||
emit: impl FnOnce(),
|
||||
) -> Vec<Value> {
|
||||
let writer = SharedWriter::default();
|
||||
let config = ObservabilityConfig::new(
|
||||
ServiceIdentity::try_new(service, "0.3.1", "test").expect("valid test identity"),
|
||||
"info",
|
||||
limits,
|
||||
);
|
||||
let subscriber =
|
||||
build_subscriber(config, writer.clone()).expect("test subscriber must be built");
|
||||
|
||||
tracing::subscriber::with_default(subscriber, emit);
|
||||
|
||||
writer
|
||||
.output()
|
||||
.lines()
|
||||
.map(|line| {
|
||||
assert!(!line.contains('\u{1b}'), "ANSI is forbidden: {line}");
|
||||
serde_json::from_str(line).expect("every line must be one JSON object")
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn schema_contract_is_identical_for_both_services() {
|
||||
let outputs = ["admin-api", "mcp-server"].map(|service| {
|
||||
capture(service, || {
|
||||
tracing::info!(
|
||||
name: "service.started",
|
||||
target: "crank::startup",
|
||||
port = 3101_u64,
|
||||
"service started"
|
||||
);
|
||||
})
|
||||
});
|
||||
|
||||
for (service, events) in ["admin-api", "mcp-server"].into_iter().zip(outputs.iter()) {
|
||||
assert_eq!(events.len(), 1);
|
||||
let event = &events[0];
|
||||
assert_eq!(event["service"], service);
|
||||
assert_eq!(event["version"], "0.3.1");
|
||||
assert_eq!(event["environment"], "test");
|
||||
assert_eq!(event["level"], "INFO");
|
||||
assert_eq!(event["target"], "crank::startup");
|
||||
assert_eq!(event["event"], "service.started");
|
||||
assert!(event["fields"].is_object());
|
||||
assert_eq!(event["fields"]["port"], 3101);
|
||||
assert_eq!(event["fields"]["message"], "service started");
|
||||
let timestamp = event["timestamp"].as_str().expect("timestamp string");
|
||||
let parsed =
|
||||
OffsetDateTime::parse(timestamp, &Rfc3339).expect("timestamp must be RFC 3339");
|
||||
assert_eq!(parsed.offset(), time::UtcOffset::UTC);
|
||||
}
|
||||
|
||||
let first_keys: Vec<_> = outputs[0][0]
|
||||
.as_object()
|
||||
.expect("object")
|
||||
.keys()
|
||||
.cloned()
|
||||
.collect();
|
||||
let second_keys: Vec<_> = outputs[1][0]
|
||||
.as_object()
|
||||
.expect("object")
|
||||
.keys()
|
||||
.cloned()
|
||||
.collect();
|
||||
assert_eq!(first_keys, second_keys);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn correlation_fields_are_distinct_and_only_present_when_recorded() {
|
||||
let present = capture("admin-api", || {
|
||||
tracing::info!(
|
||||
name: "admin.request.completed",
|
||||
request_id = "req-123",
|
||||
trace_id = "trace-456"
|
||||
);
|
||||
});
|
||||
assert_eq!(present[0]["request_id"], "req-123");
|
||||
assert_eq!(present[0]["trace_id"], "trace-456");
|
||||
assert!(present[0]["fields"].get("request_id").is_none());
|
||||
assert!(present[0]["fields"].get("trace_id").is_none());
|
||||
|
||||
let absent = capture("mcp-server", || {
|
||||
tracing::info!(name: "mcp.request.completed", status = 200_u64);
|
||||
});
|
||||
assert!(absent[0].get("request_id").is_none());
|
||||
assert!(absent[0].get("trace_id").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn correlation_fields_preserve_scalar_display_values_before_field_limits() {
|
||||
let limits = RedactionLimits {
|
||||
max_object_fields: 1,
|
||||
..RedactionLimits::default()
|
||||
};
|
||||
let request_id = "123";
|
||||
let events = capture_with_limits("admin-api", limits, || {
|
||||
tracing::info!(
|
||||
name: "admin.request.completed",
|
||||
alpha = "field that consumes the object budget",
|
||||
request_id = %request_id,
|
||||
trace_id = true,
|
||||
);
|
||||
});
|
||||
|
||||
assert_eq!(events[0]["request_id"], "123");
|
||||
assert_eq!(events[0]["trace_id"], "true");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn empty_correlation_fields_are_omitted() {
|
||||
let events = capture("admin-api", || {
|
||||
tracing::info!(
|
||||
name: "admin.request.completed",
|
||||
request_id = "",
|
||||
trace_id = ""
|
||||
);
|
||||
});
|
||||
|
||||
assert!(events[0].get("request_id").is_none());
|
||||
assert!(events[0].get("trace_id").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn formatter_redacts_fields_before_serialization() {
|
||||
let context = safe_json(
|
||||
&serde_json::json!({
|
||||
"nested": {
|
||||
"access_token": "nested-canary-secret",
|
||||
"endpoint": "https://example.test/private?key=nested-canary-secret"
|
||||
}
|
||||
}),
|
||||
RedactionLimits::default(),
|
||||
)
|
||||
.expect("safe nested context");
|
||||
let events = capture("admin-api", || {
|
||||
tracing::warn!(
|
||||
name: "admin.request.rejected",
|
||||
password = "canary-secret",
|
||||
url = "https://example.test/path?token=canary-secret",
|
||||
safe_fields = %context,
|
||||
unsafe_context = ?serde_json::json!({"password": "debug-canary-secret"}),
|
||||
error_code = "invalid_request"
|
||||
);
|
||||
});
|
||||
let serialized = serde_json::to_string(&events[0]).expect("event JSON");
|
||||
|
||||
assert_eq!(events[0]["fields"]["password"], "[REDACTED]");
|
||||
assert_eq!(events[0]["fields"]["url"], "https://example.test/path");
|
||||
assert_eq!(
|
||||
events[0]["fields"]["safe_fields"]["nested"]["access_token"],
|
||||
"[REDACTED]"
|
||||
);
|
||||
assert_eq!(
|
||||
events[0]["fields"]["safe_fields"]["nested"]["endpoint"],
|
||||
"https://example.test/private"
|
||||
);
|
||||
assert_eq!(events[0]["fields"]["unsafe_context"], "[REDACTED]");
|
||||
assert_eq!(events[0]["fields"]["error_code"], "invalid_request");
|
||||
assert!(!serialized.contains("canary-secret"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn arbitrary_debug_text_is_never_written_verbatim() {
|
||||
struct Credentials;
|
||||
|
||||
impl std::fmt::Debug for Credentials {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter.write_str("Credentials { password: \"debug-canary-secret\" }")
|
||||
}
|
||||
}
|
||||
|
||||
let events = capture("admin-api", || {
|
||||
tracing::warn!(
|
||||
name: "admin.debug.inspected",
|
||||
details = ?Credentials
|
||||
);
|
||||
});
|
||||
let serialized = serde_json::to_string(&events[0]).expect("event JSON");
|
||||
|
||||
assert_eq!(events[0]["fields"]["details"], "[REDACTED]");
|
||||
assert!(!serialized.contains("debug-canary-secret"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn compound_sensitive_event_fields_are_redacted() {
|
||||
let events = capture("admin-api", || {
|
||||
tracing::warn!(
|
||||
name: "admin.request.rejected",
|
||||
client_api_key = "client-canary-secret",
|
||||
authorization_header = "Bearer auth-canary-secret",
|
||||
response_body = "response-canary-secret",
|
||||
tool_arguments = "argument-canary-secret",
|
||||
query_params = "query-canary-secret",
|
||||
);
|
||||
});
|
||||
let serialized = serde_json::to_string(&events[0]).expect("event JSON");
|
||||
|
||||
for key in [
|
||||
"client_api_key",
|
||||
"authorization_header",
|
||||
"response_body",
|
||||
"tool_arguments",
|
||||
"query_params",
|
||||
] {
|
||||
assert_eq!(events[0]["fields"][key], "[REDACTED]");
|
||||
}
|
||||
assert!(!serialized.contains("canary-secret"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn oversized_event_falls_back_to_valid_bounded_json() {
|
||||
let limits = RedactionLimits {
|
||||
max_event_bytes: 512,
|
||||
..RedactionLimits::default()
|
||||
};
|
||||
let events = capture_with_limits("admin-api", limits, || {
|
||||
tracing::info!(
|
||||
name: "admin.payload.inspected",
|
||||
description = %"x".repeat(1024)
|
||||
);
|
||||
});
|
||||
let serialized = serde_json::to_vec(&events[0]).expect("bounded event JSON");
|
||||
|
||||
assert!(serialized.len() < limits.max_event_bytes);
|
||||
assert_eq!(events[0]["fields"]["truncated"], true);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn safe_json_honours_the_total_event_budget() {
|
||||
let limits = RedactionLimits {
|
||||
max_event_bytes: 512,
|
||||
..RedactionLimits::default()
|
||||
};
|
||||
|
||||
let serialized = safe_json(
|
||||
&serde_json::json!({"description": "x".repeat(4096)}),
|
||||
limits,
|
||||
)
|
||||
.expect("safe JSON must remain serializable");
|
||||
|
||||
assert!(serialized.len() <= limits.max_event_bytes);
|
||||
assert_eq!(
|
||||
serde_json::from_str::<Value>(&serialized).expect("valid JSON")["truncated"],
|
||||
true
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn subscriber_rejects_limits_that_cannot_hold_an_event() {
|
||||
let config = ObservabilityConfig::new(
|
||||
ServiceIdentity::try_new("admin-api", "0.3.1", "test").expect("valid identity"),
|
||||
"info",
|
||||
RedactionLimits {
|
||||
max_event_bytes: 16,
|
||||
..RedactionLimits::default()
|
||||
},
|
||||
);
|
||||
|
||||
assert!(build_subscriber(config, SharedWriter::default()).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn minimum_event_budget_handles_maximum_identity_labels() {
|
||||
let writer = SharedWriter::default();
|
||||
let config = ObservabilityConfig::new(
|
||||
ServiceIdentity::try_new("s".repeat(64), "v".repeat(64), "e".repeat(64))
|
||||
.expect("maximum identity labels are valid"),
|
||||
"info",
|
||||
RedactionLimits {
|
||||
max_event_bytes: 512,
|
||||
..RedactionLimits::default()
|
||||
},
|
||||
);
|
||||
let subscriber =
|
||||
build_subscriber(config, writer.clone()).expect("minimum valid budget must be usable");
|
||||
|
||||
tracing::subscriber::with_default(subscriber, || {
|
||||
tracing::info!(
|
||||
name: "event-name-that-is-intentionally-longer-than-the-fallback-limit",
|
||||
description = %"x".repeat(4096),
|
||||
);
|
||||
});
|
||||
|
||||
let output = writer.output();
|
||||
assert!(output.len() <= 512);
|
||||
assert_eq!(output.lines().count(), 1);
|
||||
serde_json::from_str::<Value>(output.trim_end()).expect("bounded line must remain valid JSON");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn env_filter_is_applied_and_invalid_filter_is_safe() {
|
||||
let writer = SharedWriter::default();
|
||||
let config = ObservabilityConfig::new(
|
||||
ServiceIdentity::try_new("admin-api", "0.3.1", "test").expect("valid identity"),
|
||||
"warn",
|
||||
RedactionLimits::default(),
|
||||
);
|
||||
let subscriber =
|
||||
build_subscriber(config, writer.clone()).expect("test subscriber must be built");
|
||||
tracing::subscriber::with_default(subscriber, || {
|
||||
tracing::info!(name: "filtered.info", "filtered");
|
||||
tracing::warn!(name: "visible.warning", "visible");
|
||||
});
|
||||
let output = writer.output();
|
||||
|
||||
assert!(!output.contains("filtered.info"));
|
||||
assert!(output.contains("visible.warning"));
|
||||
|
||||
let invalid = ObservabilityConfig::new(
|
||||
ServiceIdentity::try_new("admin-api", "0.3.1", "test").expect("valid identity"),
|
||||
"[not a valid filter",
|
||||
RedactionLimits::default(),
|
||||
);
|
||||
let error = build_subscriber(invalid, SharedWriter::default())
|
||||
.err()
|
||||
.expect("invalid filter must fail");
|
||||
assert_eq!(error.to_string(), "invalid log filter");
|
||||
assert!(!error.to_string().contains("not a valid filter"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn service_identity_rejects_empty_or_unsafe_labels() {
|
||||
for (service, version, environment) in [
|
||||
("", "0.3.1", "test"),
|
||||
("admin api", "0.3.1", "test"),
|
||||
("admin-api", "", "test"),
|
||||
("admin-api", "0.3.1", "prod\nsecret"),
|
||||
] {
|
||||
assert!(ServiceIdentity::try_new(service, version, environment).is_err());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
use crank_observability::{
|
||||
ObservabilityConfig, ObservabilityInitError, RedactionLimits, ServiceIdentity, init,
|
||||
};
|
||||
|
||||
fn config() -> ObservabilityConfig {
|
||||
ObservabilityConfig::new(
|
||||
ServiceIdentity::try_new("lifecycle-test", "0.3.1", "test").expect("valid test identity"),
|
||||
"info",
|
||||
RedactionLimits::default(),
|
||||
)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn repeated_global_initialization_returns_typed_error() {
|
||||
let _lifecycle = init(config()).expect("first initialization must succeed");
|
||||
let error = init(config()).expect_err("second initialization must fail");
|
||||
|
||||
assert!(matches!(
|
||||
error,
|
||||
ObservabilityInitError::SubscriberAlreadyInitialized
|
||||
));
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
use std::time::Duration;
|
||||
|
||||
use crank_observability::{OtlpBatchConfig, OtlpTraceConfig, OtlpTraceConfigError};
|
||||
|
||||
#[test]
|
||||
fn absent_endpoint_disables_export_without_background_resources() {
|
||||
let config = OtlpTraceConfig::try_new(
|
||||
None,
|
||||
None,
|
||||
Duration::from_secs(10),
|
||||
OtlpBatchConfig::default(),
|
||||
)
|
||||
.expect("missing endpoint must be valid");
|
||||
|
||||
assert!(!config.is_enabled());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn explicit_config_accepts_only_bounded_http_protobuf() {
|
||||
let config = OtlpTraceConfig::try_new(
|
||||
Some("https://collector.example.test/v1/traces".to_owned()),
|
||||
Some("http/protobuf".to_owned()),
|
||||
Duration::from_secs(3),
|
||||
OtlpBatchConfig::try_new(256, 64, Duration::from_millis(500), Duration::from_secs(3))
|
||||
.unwrap(),
|
||||
)
|
||||
.expect("bounded HTTP protobuf config must be valid");
|
||||
|
||||
assert!(config.is_enabled());
|
||||
assert_eq!(config.export_timeout(), Duration::from_secs(3));
|
||||
assert_eq!(config.batch().max_queue_size(), 256);
|
||||
assert_eq!(config.batch().max_export_batch_size(), 64);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_values_return_safe_typed_errors() {
|
||||
let secret_endpoint = "https://user:canary-secret@collector.example.test/v1/traces";
|
||||
let error = OtlpTraceConfig::try_new(
|
||||
Some(secret_endpoint.to_owned()),
|
||||
Some("grpc".to_owned()),
|
||||
Duration::ZERO,
|
||||
OtlpBatchConfig::default(),
|
||||
)
|
||||
.expect_err("credentials in endpoint must be rejected");
|
||||
|
||||
assert!(matches!(
|
||||
error,
|
||||
OtlpTraceConfigError::InvalidEndpoint { .. }
|
||||
));
|
||||
assert!(!error.to_string().contains(secret_endpoint));
|
||||
assert!(!error.to_string().contains("canary-secret"));
|
||||
|
||||
let error = OtlpBatchConfig::try_new(8, 9, Duration::from_millis(1), Duration::from_secs(1))
|
||||
.expect_err("batch cannot exceed queue");
|
||||
assert!(matches!(error, OtlpTraceConfigError::InvalidBatchLimits));
|
||||
}
|
||||
@@ -0,0 +1,134 @@
|
||||
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
|
||||
|
||||
use axum::{
|
||||
body::{Body, to_bytes},
|
||||
http::{Request, StatusCode, header},
|
||||
};
|
||||
use crank_observability::{
|
||||
DURATION_BUCKETS_SECONDS, MetricsConfig, MetricsConfigError, MetricsSurface, ServiceIdentity,
|
||||
metric_schema,
|
||||
};
|
||||
use tower::ServiceExt;
|
||||
|
||||
fn identity() -> ServiceIdentity {
|
||||
ServiceIdentity::try_new("admin-api", "0.3.1", "test").expect("valid identity")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn loopback_is_allowed_without_a_token() {
|
||||
let config = MetricsConfig::new(
|
||||
true,
|
||||
SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 9464),
|
||||
None,
|
||||
)
|
||||
.expect("loopback metrics must be safe by default");
|
||||
|
||||
assert_eq!(config.bind_addr().to_string(), "127.0.0.1:9464");
|
||||
assert!(!config.requires_authentication());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn non_loopback_without_a_token_is_rejected_without_secret_data() {
|
||||
let error = MetricsConfig::new(
|
||||
true,
|
||||
SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 9464),
|
||||
None,
|
||||
)
|
||||
.expect_err("external metrics must require authentication");
|
||||
|
||||
assert!(matches!(
|
||||
error,
|
||||
MetricsConfigError::MissingTokenForExternalBind
|
||||
));
|
||||
assert!(!error.to_string().contains("token="));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn schema_is_closed_and_uses_fixed_duration_buckets() {
|
||||
let schema = metric_schema();
|
||||
let names: Vec<_> = schema.iter().map(|metric| metric.name).collect();
|
||||
|
||||
assert!(names.contains(&"crank_http_requests_total"));
|
||||
assert!(names.contains(&"crank_http_request_duration_seconds"));
|
||||
assert!(names.contains(&"crank_mcp_requests_total"));
|
||||
assert!(names.contains(&"crank_tool_invocations_total"));
|
||||
assert!(names.contains(&"crank_runtime_inflight"));
|
||||
assert!(names.contains(&"crank_db_pool_connections"));
|
||||
assert!(names.contains(&"crank_catalog_tools"));
|
||||
assert!(names.contains(&"crank_invocation_history_lost_total"));
|
||||
assert!(names.contains(&"crank_telemetry_export_failures_total"));
|
||||
|
||||
for metric in schema {
|
||||
for forbidden in [
|
||||
"workspace",
|
||||
"agent_id",
|
||||
"operation_id",
|
||||
"request_id",
|
||||
"url",
|
||||
"error_message",
|
||||
"text",
|
||||
] {
|
||||
assert!(
|
||||
!metric.labels.contains(&forbidden),
|
||||
"{} exposes forbidden label {forbidden}",
|
||||
metric.name
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
assert_eq!(
|
||||
DURATION_BUCKETS_SECONDS,
|
||||
&[
|
||||
0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0, 30.0, 60.0
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn external_surface_protects_both_routes_and_exposes_nothing_else() {
|
||||
let token = "metrics-canary-secret";
|
||||
let config = MetricsConfig::new(
|
||||
true,
|
||||
SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 9464),
|
||||
Some(token.to_owned()),
|
||||
)
|
||||
.expect("external metrics with token");
|
||||
let surface = MetricsSurface::for_test(config, identity()).expect("test metrics surface");
|
||||
let app = surface.router();
|
||||
|
||||
for path in ["/metrics", "/health"] {
|
||||
let unauthorized = app
|
||||
.clone()
|
||||
.oneshot(Request::get(path).body(Body::empty()).expect("request"))
|
||||
.await
|
||||
.expect("response");
|
||||
assert_eq!(unauthorized.status(), StatusCode::UNAUTHORIZED);
|
||||
|
||||
let authorized = app
|
||||
.clone()
|
||||
.oneshot(
|
||||
Request::get(path)
|
||||
.header(header::AUTHORIZATION, format!("Bearer {token}"))
|
||||
.body(Body::empty())
|
||||
.expect("request"),
|
||||
)
|
||||
.await
|
||||
.expect("response");
|
||||
assert_eq!(authorized.status(), StatusCode::OK);
|
||||
let body = to_bytes(authorized.into_body(), 1024 * 1024)
|
||||
.await
|
||||
.expect("bounded body");
|
||||
assert!(!String::from_utf8_lossy(&body).contains(token));
|
||||
}
|
||||
|
||||
let absent = app
|
||||
.oneshot(
|
||||
Request::get("/api/operations")
|
||||
.header(header::AUTHORIZATION, format!("Bearer {token}"))
|
||||
.body(Body::empty())
|
||||
.expect("request"),
|
||||
)
|
||||
.await
|
||||
.expect("response");
|
||||
assert_eq!(absent.status(), StatusCode::NOT_FOUND);
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
use axum::http::{HeaderMap, HeaderValue};
|
||||
use crank_observability::{inject_current_trace_context, set_remote_trace_parent};
|
||||
use opentelemetry::{global, trace::TracerProvider as _};
|
||||
use opentelemetry_sdk::{propagation::TraceContextPropagator, trace::SdkTracerProvider};
|
||||
use tracing::info_span;
|
||||
use tracing_subscriber::layer::SubscriberExt;
|
||||
|
||||
const REMOTE_TRACE_ID: &str = "0af7651916cd43dd8448eb211c80319c";
|
||||
|
||||
#[test]
|
||||
fn valid_remote_parent_is_continued_and_request_id_is_unrelated() {
|
||||
with_trace_dispatch(|| {
|
||||
let mut incoming = HeaderMap::new();
|
||||
incoming.insert(
|
||||
"traceparent",
|
||||
HeaderValue::from_static("00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01"),
|
||||
);
|
||||
incoming.insert("x-request-id", HeaderValue::from_static("req-unrelated"));
|
||||
|
||||
let span = info_span!("http.request", request_id = "req-unrelated");
|
||||
assert!(set_remote_trace_parent(&span, &incoming));
|
||||
let _guard = span.enter();
|
||||
|
||||
let mut outgoing = HeaderMap::new();
|
||||
assert!(inject_current_trace_context(&mut outgoing));
|
||||
let propagated = outgoing["traceparent"].to_str().unwrap();
|
||||
assert_eq!(&propagated[3..35], REMOTE_TRACE_ID);
|
||||
assert!(!propagated.contains("req-unrelated"));
|
||||
assert!(!outgoing.contains_key("baggage"));
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_parent_is_ignored_and_a_new_trace_is_created() {
|
||||
with_trace_dispatch(|| {
|
||||
let mut incoming = HeaderMap::new();
|
||||
incoming.insert(
|
||||
"traceparent",
|
||||
HeaderValue::from_static("canary-invalid-traceparent"),
|
||||
);
|
||||
|
||||
let span = info_span!("mcp.request");
|
||||
assert!(!set_remote_trace_parent(&span, &incoming));
|
||||
let _guard = span.enter();
|
||||
|
||||
let mut outgoing = HeaderMap::new();
|
||||
assert!(inject_current_trace_context(&mut outgoing));
|
||||
let propagated = outgoing["traceparent"].to_str().unwrap();
|
||||
assert!(propagated.starts_with("00-"));
|
||||
assert!(!propagated.contains("canary-invalid-traceparent"));
|
||||
});
|
||||
}
|
||||
|
||||
fn with_trace_dispatch(test: impl FnOnce()) {
|
||||
global::set_text_map_propagator(TraceContextPropagator::new());
|
||||
let provider = SdkTracerProvider::builder().build();
|
||||
let tracer = provider.tracer("propagation-test");
|
||||
let subscriber =
|
||||
tracing_subscriber::registry().with(tracing_opentelemetry::layer().with_tracer(tracer));
|
||||
let dispatch = tracing::Dispatch::new(subscriber);
|
||||
let _guard = tracing::dispatcher::set_default(&dispatch);
|
||||
|
||||
test();
|
||||
provider.shutdown().unwrap();
|
||||
}
|
||||
@@ -0,0 +1,206 @@
|
||||
use crank_observability::{RedactionLimits, redact_value};
|
||||
use serde_json::{Value, json};
|
||||
|
||||
const REDACTED: &str = "[REDACTED]";
|
||||
const TRUNCATED: &str = "[TRUNCATED]";
|
||||
|
||||
#[test]
|
||||
fn sensitive_keys_are_redacted_case_insensitively() {
|
||||
let keys = [
|
||||
"password",
|
||||
"PassWord",
|
||||
"api_key",
|
||||
"access-token",
|
||||
"authorization",
|
||||
"Proxy.Authorization",
|
||||
"cookie",
|
||||
"set_cookie",
|
||||
"payload",
|
||||
"request_body",
|
||||
"arguments",
|
||||
"result",
|
||||
"response",
|
||||
];
|
||||
|
||||
for key in keys {
|
||||
let cleaned = redact_value(&json!({ key: "canary-secret" }), RedactionLimits::default());
|
||||
assert_eq!(cleaned[key], REDACTED, "key {key} was not redacted");
|
||||
assert!(!cleaned.to_string().contains("canary-secret"));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generated_sensitive_key_variants_are_redacted() {
|
||||
for canonical in [
|
||||
"password",
|
||||
"api_key",
|
||||
"access_key",
|
||||
"secret_key",
|
||||
"proxy_authorization",
|
||||
"set_cookie",
|
||||
"query_string",
|
||||
"request_body",
|
||||
] {
|
||||
for separator in ["_", "-", "."] {
|
||||
let variant = canonical
|
||||
.split('_')
|
||||
.collect::<Vec<_>>()
|
||||
.join(separator)
|
||||
.to_ascii_uppercase();
|
||||
let cleaned = redact_value(
|
||||
&json!({ variant.clone(): "canary-secret" }),
|
||||
RedactionLimits::default(),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
cleaned[&variant], REDACTED,
|
||||
"key {variant} was not redacted"
|
||||
);
|
||||
assert!(!cleaned.to_string().contains("canary-secret"));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn compound_sensitive_key_names_are_redacted() {
|
||||
for key in [
|
||||
"client_api_key",
|
||||
"aws_access_key_id",
|
||||
"http_authorization_header",
|
||||
"response_body",
|
||||
"tool_arguments",
|
||||
"query_params",
|
||||
"secret_value",
|
||||
] {
|
||||
let cleaned = redact_value(&json!({ key: "canary-secret" }), RedactionLimits::default());
|
||||
|
||||
assert_eq!(cleaned[key], REDACTED, "key {key} was not redacted");
|
||||
assert!(!cleaned.to_string().contains("canary-secret"));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn url_query_and_fragment_are_removed_at_every_depth() {
|
||||
let input = json!({
|
||||
"url": "https://example.test/path?token=canary-secret#fragment",
|
||||
"nested": [{
|
||||
"endpoint_uri": "https://example.test/other?q=canary-secret"
|
||||
}]
|
||||
});
|
||||
|
||||
let cleaned = redact_value(&input, RedactionLimits::default());
|
||||
|
||||
assert_eq!(cleaned["url"], "https://example.test/path");
|
||||
assert_eq!(
|
||||
cleaned["nested"][0]["endpoint_uri"],
|
||||
"https://example.test/other"
|
||||
);
|
||||
assert!(!cleaned.to_string().contains("canary-secret"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn url_credentials_and_compound_endpoint_fields_are_removed() {
|
||||
let input = json!({
|
||||
"upstream_endpoint": "https://user:canary-secret@example.test/path?token=canary-secret",
|
||||
});
|
||||
|
||||
let cleaned = redact_value(&input, RedactionLimits::default());
|
||||
|
||||
assert_eq!(cleaned["upstream_endpoint"], "https://example.test/path");
|
||||
assert!(!cleaned.to_string().contains("user"));
|
||||
assert!(!cleaned.to_string().contains("canary-secret"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn nested_values_and_collections_respect_all_limits() {
|
||||
let limits = RedactionLimits {
|
||||
max_string_bytes: 16,
|
||||
max_array_items: 3,
|
||||
max_object_fields: 3,
|
||||
max_depth: 2,
|
||||
max_event_bytes: 256,
|
||||
};
|
||||
let input = json!({
|
||||
"long": "абвгдежзийклмнопрсту",
|
||||
"array": [1, 2, 3, 4, 5],
|
||||
"object": {"a": 1, "b": 2, "c": 3, "d": 4},
|
||||
"nested": {"level2": {"level3": "must not survive"}}
|
||||
});
|
||||
|
||||
let cleaned = redact_value(&input, limits);
|
||||
let object = cleaned
|
||||
.as_object()
|
||||
.expect("cleaned root must remain an object");
|
||||
|
||||
assert!(object.len() <= limits.max_object_fields);
|
||||
assert!(
|
||||
cleaned["long"]
|
||||
.as_str()
|
||||
.map(|value| value.len() <= limits.max_string_bytes)
|
||||
.unwrap_or(true)
|
||||
);
|
||||
assert!(
|
||||
cleaned["array"]
|
||||
.as_array()
|
||||
.map(|value| value.len() <= limits.max_array_items)
|
||||
.unwrap_or(true)
|
||||
);
|
||||
assert!(!cleaned.to_string().contains("must not survive"));
|
||||
assert!(cleaned.to_string().contains(TRUNCATED));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn truncation_preserves_utf8_and_never_reveals_secret_fragments() {
|
||||
let input = json!({
|
||||
"secret_key": "секретное-значение",
|
||||
"description": "я".repeat(2048),
|
||||
});
|
||||
|
||||
let cleaned = redact_value(&input, RedactionLimits::default());
|
||||
let serialized = serde_json::to_string(&cleaned).expect("cleaned value must be valid JSON");
|
||||
|
||||
assert_eq!(cleaned["secret_key"], REDACTED);
|
||||
assert!(!serialized.contains("секретное"));
|
||||
assert!(cleaned["description"].as_str().is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn redacted_value_does_not_mutate_input() {
|
||||
let input = json!({"password": "canary-secret"});
|
||||
let original = input.clone();
|
||||
|
||||
let _ = redact_value(&input, RedactionLimits::default());
|
||||
|
||||
assert_eq!(input, original);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn object_keys_respect_the_string_limit() {
|
||||
let limits = RedactionLimits {
|
||||
max_string_bytes: 16,
|
||||
..RedactionLimits::default()
|
||||
};
|
||||
let long_key = format!("field-{}", "x".repeat(128));
|
||||
|
||||
let cleaned = redact_value(&json!({ long_key: "value" }), limits);
|
||||
|
||||
assert!(
|
||||
cleaned
|
||||
.as_object()
|
||||
.expect("cleaned object")
|
||||
.keys()
|
||||
.all(|key| key.len() <= limits.max_string_bytes)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn limits_have_finite_safe_defaults() {
|
||||
let limits = RedactionLimits::default();
|
||||
|
||||
assert_eq!(limits.max_string_bytes, 1024);
|
||||
assert_eq!(limits.max_array_items, 32);
|
||||
assert_eq!(limits.max_object_fields, 64);
|
||||
assert_eq!(limits.max_depth, 8);
|
||||
assert_eq!(limits.max_event_bytes, 16 * 1024);
|
||||
assert!(Value::Null.is_null());
|
||||
}
|
||||
@@ -3,6 +3,7 @@ name = "crank-registry"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
rust-version.workspace = true
|
||||
publish.workspace = true
|
||||
version.workspace = true
|
||||
|
||||
[dependencies]
|
||||
|
||||
@@ -74,6 +74,8 @@ pub enum RegistryError {
|
||||
YamlImportJobNotFound { job_id: String },
|
||||
#[error("import job {job_id} was not found")]
|
||||
ImportJobNotFound { job_id: String },
|
||||
#[error("import job {job_id} was already applied with different parameters")]
|
||||
ImportJobAlreadyApplied { job_id: String },
|
||||
#[error("unsupported enum representation for field {field}")]
|
||||
InvalidEnumRepresentation { field: &'static str },
|
||||
#[error("invalid numeric value for field {field}: {value}")]
|
||||
|
||||
@@ -9,12 +9,14 @@ pub use ext::{ExtensionMigration, RegistryExtension, apply_extension_migrations}
|
||||
|
||||
pub mod records {
|
||||
pub use crate::model::{
|
||||
AgentSummary, AgentVersionRecord, ApprovalRequestRecord, AuthUserRecord, DescriptorKind,
|
||||
DescriptorMetadata, ImportJob, ImportJobId, ImportJobKind, ImportJobStatus,
|
||||
InvitationRecord, InvocationLogRecord, MembershipRecord, OperationAgentRef,
|
||||
OperationSampleMetadata, OperationSummary, OperationUsageSummary, OperationVersionRecord,
|
||||
Page, PlatformApiKeyRecord, PublishedAgentCatalog, PublishedAgentTool, RegistryOperation,
|
||||
SampleKind, SecretRecord, SecretVersionRecord, SessionRecord, UsageAgentBreakdown,
|
||||
AgentSummary, AgentVersionRecord, AppliedImportOperation, ApprovalRequestRecord,
|
||||
AuthUserRecord, DescriptorKind, DescriptorMetadata, ImportJob, ImportJobApplyResult,
|
||||
ImportJobId, ImportJobKind, ImportJobStatus, InvitationRecord, InvocationHistoryLoss,
|
||||
InvocationHistoryLossCategory, InvocationHistoryWriteOutcome, InvocationLogRecord,
|
||||
MembershipRecord, OperationAgentRef, OperationSampleMetadata, OperationSummary,
|
||||
OperationUsageSummary, OperationVersionRecord, Page, PlatformApiKeyRecord,
|
||||
PublishedAgentCatalog, PublishedAgentTool, RegistryOperation, SampleKind, SecretRecord,
|
||||
SecretVersionRecord, SessionRecord, SkippedImportOperation, UsageAgentBreakdown,
|
||||
UsageBucket, UsageOperationBreakdown, UsageRollupRecord, UsageSummary, UsageTimelinePoint,
|
||||
WorkspaceMembershipRecord, WorkspaceRecord, WorkspaceUpstream, WorkspaceUpstreamId,
|
||||
YamlImportJob, YamlImportJobCompletion, YamlImportJobId, YamlImportJobStatus,
|
||||
@@ -23,11 +25,12 @@ pub mod records {
|
||||
|
||||
pub mod requests {
|
||||
pub use crate::model::{
|
||||
CreateAgentDraftVersionRequest, CreateAgentRequest, CreateApprovalRequest,
|
||||
CreateImportJobRequest, CreateInvitationRequest, CreateInvocationLogRequest,
|
||||
CreatePlatformApiKeyRequest, CreateSecretRequest, CreateVersionRequest,
|
||||
CreateWorkspaceRequest, CreateYamlImportJobRequest, DecideApprovalRequest,
|
||||
ExpireApprovalRequest, FinishApprovalRequest, FinishImportJobRequest,
|
||||
ApplyImportJobRequest, CreateAgentDraftVersionRequest, CreateAgentRequest,
|
||||
CreateApprovalRequest, CreateImportJobRequest, CreateInvitationRequest,
|
||||
CreateInvocationLogRequest, CreatePlatformApiKeyRequest, CreateSecretRequest,
|
||||
CreateVersionRequest, CreateWorkspaceRequest, CreateYamlImportJobRequest,
|
||||
DecideApprovalRequest, ExpireApprovalRequest, FinishApprovalRequest,
|
||||
FinishImportJobRequest, ImportConflictMode, ImportOperationDraft,
|
||||
ListApprovalRequestsQuery, ListInvocationLogsQuery, PublishAgentRequest, PublishRequest,
|
||||
RotateSecretRequest, SaveAgentBindingsRequest, SaveAgentCatalogConfigRequest,
|
||||
SaveAuthProfileRequest, SaveDescriptorMetadataRequest, SaveSampleMetadataRequest,
|
||||
@@ -41,20 +44,23 @@ pub mod infrastructure {
|
||||
}
|
||||
|
||||
pub use model::{
|
||||
AgentSummary, AgentVersionRecord, ApprovalRequestRecord, AuthUserRecord,
|
||||
CreateAgentDraftVersionRequest, CreateAgentRequest, CreateApprovalRequest,
|
||||
CreateImportJobRequest, CreateInvitationRequest, CreateInvocationLogRequest,
|
||||
CreatePlatformApiKeyRequest, CreateSecretRequest, CreateVersionRequest, CreateWorkspaceRequest,
|
||||
CreateYamlImportJobRequest, DecideApprovalRequest, DescriptorKind, DescriptorMetadata,
|
||||
ExpireApprovalRequest, FinishApprovalRequest, FinishImportJobRequest, ImportJob, ImportJobId,
|
||||
ImportJobKind, ImportJobStatus, InvitationRecord, InvocationLogRecord,
|
||||
ListApprovalRequestsQuery, ListInvocationLogsQuery, MembershipRecord, OperationAgentRef,
|
||||
OperationSampleMetadata, OperationSummary, OperationUsageSummary, OperationVersionRecord, Page,
|
||||
PlatformApiKeyRecord, PublishAgentRequest, PublishRequest, PublishedAgentCatalog,
|
||||
PublishedAgentTool, RegistryOperation, RotateSecretRequest, SampleKind,
|
||||
SaveAgentBindingsRequest, SaveAgentCatalogConfigRequest, SaveAuthProfileRequest,
|
||||
SaveDescriptorMetadataRequest, SaveSampleMetadataRequest, SaveWorkspaceUpstreamRequest,
|
||||
SecretRecord, SecretVersionRecord, SessionRecord, UpdateWorkspaceRequest, UsageAgentBreakdown,
|
||||
AgentSummary, AgentVersionRecord, AppliedImportOperation, ApplyImportJobRequest,
|
||||
ApprovalRequestRecord, AuthUserRecord, CreateAgentDraftVersionRequest, CreateAgentRequest,
|
||||
CreateApprovalRequest, CreateImportJobRequest, CreateInvitationRequest,
|
||||
CreateInvocationLogRequest, CreatePlatformApiKeyRequest, CreateSecretRequest,
|
||||
CreateVersionRequest, CreateWorkspaceRequest, CreateYamlImportJobRequest,
|
||||
DecideApprovalRequest, DescriptorKind, DescriptorMetadata, ExpireApprovalRequest,
|
||||
FinishApprovalRequest, FinishImportJobRequest, ImportConflictMode, ImportJob,
|
||||
ImportJobApplyResult, ImportJobId, ImportJobKind, ImportJobStatus, ImportOperationDraft,
|
||||
InvitationRecord, InvocationHistoryLoss, InvocationHistoryLossCategory,
|
||||
InvocationHistoryWriteOutcome, InvocationLogRecord, ListApprovalRequestsQuery,
|
||||
ListInvocationLogsQuery, MembershipRecord, OperationAgentRef, OperationSampleMetadata,
|
||||
OperationSummary, OperationUsageSummary, OperationVersionRecord, Page, PlatformApiKeyRecord,
|
||||
PublishAgentRequest, PublishRequest, PublishedAgentCatalog, PublishedAgentTool,
|
||||
RegistryOperation, RotateSecretRequest, SampleKind, SaveAgentBindingsRequest,
|
||||
SaveAgentCatalogConfigRequest, SaveAuthProfileRequest, SaveDescriptorMetadataRequest,
|
||||
SaveSampleMetadataRequest, SaveWorkspaceUpstreamRequest, SecretRecord, SecretVersionRecord,
|
||||
SessionRecord, SkippedImportOperation, UpdateWorkspaceRequest, UsageAgentBreakdown,
|
||||
UsageBucket, UsageOperationBreakdown, UsageQuery, UsageRollupRecord, UsageSummary,
|
||||
UsageTimelinePoint, WorkspaceMembershipRecord, WorkspaceRecord, WorkspaceUpstream,
|
||||
WorkspaceUpstreamId, YamlImportJob, YamlImportJobCompletion, YamlImportJobId,
|
||||
|
||||
@@ -1,6 +1,61 @@
|
||||
use sqlx::{PgPool, query};
|
||||
use sqlx::{PgPool, Postgres, Row, Transaction, query};
|
||||
|
||||
const CORE_MIGRATION_LOCK_ID: i64 = 0x43_52_41_4E_4B;
|
||||
const BASELINE_VERSION: i32 = 1;
|
||||
const BASELINE_CHECKSUM: &str = "crank-community-baseline-v1";
|
||||
|
||||
pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
let mut transaction = pool.begin().await?;
|
||||
query("select pg_advisory_xact_lock($1)")
|
||||
.bind(CORE_MIGRATION_LOCK_ID)
|
||||
.execute(&mut *transaction)
|
||||
.await?;
|
||||
query(
|
||||
"create table if not exists __crank_core_migrations (
|
||||
version integer primary key,
|
||||
description text not null,
|
||||
checksum text not null,
|
||||
applied_at timestamptz not null default now()
|
||||
)",
|
||||
)
|
||||
.execute(&mut *transaction)
|
||||
.await?;
|
||||
|
||||
let applied = query(
|
||||
"select version, checksum
|
||||
from __crank_core_migrations
|
||||
order by version",
|
||||
)
|
||||
.fetch_all(&mut *transaction)
|
||||
.await?;
|
||||
for row in &applied {
|
||||
let version = row.try_get::<i32, _>("version")?;
|
||||
let checksum = row.try_get::<String, _>("checksum")?;
|
||||
if version != BASELINE_VERSION || checksum != BASELINE_CHECKSUM {
|
||||
return Err(sqlx::Error::Protocol(format!(
|
||||
"unsupported or modified core migration: version={version}, checksum={checksum}"
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
if applied.is_empty() {
|
||||
apply_baseline(&mut transaction).await?;
|
||||
query(
|
||||
"insert into __crank_core_migrations (version, description, checksum)
|
||||
values ($1, $2, $3)",
|
||||
)
|
||||
.bind(BASELINE_VERSION)
|
||||
.bind("community baseline")
|
||||
.bind(BASELINE_CHECKSUM)
|
||||
.execute(&mut *transaction)
|
||||
.await?;
|
||||
}
|
||||
|
||||
transaction.commit().await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn apply_baseline(transaction: &mut Transaction<'_, Postgres>) -> Result<(), sqlx::Error> {
|
||||
query(
|
||||
"create table if not exists workspaces (
|
||||
id text primary key,
|
||||
@@ -12,7 +67,7 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
updated_at timestamptz not null
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -25,11 +80,11 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
created_at timestamptz not null
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query("alter table users add column if not exists password_hash text null")
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -48,7 +103,7 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
)
|
||||
on conflict (id) do nothing",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -60,7 +115,7 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
primary key (workspace_id, user_id)
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -75,14 +130,14 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
created_at timestamptz not null
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
"alter table user_sessions
|
||||
add column if not exists current_workspace_id text null references workspaces(id) on delete set null",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -105,7 +160,7 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
)
|
||||
on conflict (id) do nothing",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -122,7 +177,7 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
)
|
||||
on conflict (workspace_id, user_id) do nothing",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -137,7 +192,7 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
created_at timestamptz not null
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -158,29 +213,29 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
allowed_origins_json jsonb not null default '[]'::jsonb
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
"create unique index if not exists platform_api_keys_workspace_name_idx on platform_api_keys(workspace_id, name)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query("alter table platform_api_keys add column if not exists agent_id text null")
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query(
|
||||
"alter table platform_api_keys add column if not exists key_kind text not null default 'mcp_client'",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query("alter table platform_api_keys add column if not exists expires_at timestamptz null")
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query(
|
||||
"alter table platform_api_keys add column if not exists allowed_origins_json jsonb not null default '[]'::jsonb",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -203,7 +258,7 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
)
|
||||
on conflict (id) do nothing",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -223,35 +278,35 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
published_at timestamptz null
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query("alter table operations add column if not exists workspace_id text null references workspaces(id) on delete cascade")
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query(
|
||||
"alter table operations add column if not exists category text not null default 'general'",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query(
|
||||
"alter table operations add column if not exists security_level text not null default 'standard'",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query("update operations set workspace_id = 'ws_default' where workspace_id is null")
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query("alter table operations alter column workspace_id set not null")
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query("alter table operations drop constraint if exists operations_name_key")
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query(
|
||||
"create unique index if not exists operations_workspace_name_idx on operations(workspace_id, name)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -276,11 +331,11 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
primary key (operation_id, version)
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query("alter table operation_versions add column if not exists wizard_state_json jsonb null")
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -292,7 +347,7 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
foreign key (operation_id, version) references operation_versions(operation_id, version) on delete cascade
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -308,7 +363,7 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
foreign key (operation_id, version) references operation_versions(operation_id, version) on delete cascade
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -324,7 +379,7 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
foreign key (operation_id, version) references operation_versions(operation_id, version) on delete cascade
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -338,25 +393,25 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
updated_at timestamptz not null
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query("alter table auth_profiles add column if not exists workspace_id text null references workspaces(id) on delete cascade")
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query("update auth_profiles set workspace_id = 'ws_default' where workspace_id is null")
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query("alter table auth_profiles alter column workspace_id set not null")
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query("alter table auth_profiles drop constraint if exists auth_profiles_name_key")
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query(
|
||||
"create unique index if not exists auth_profiles_workspace_name_idx on auth_profiles(workspace_id, name)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -371,17 +426,17 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
updated_at timestamptz not null
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query(
|
||||
"create unique index if not exists workspace_upstreams_workspace_name_idx on workspace_upstreams(workspace_id, name)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query(
|
||||
"create unique index if not exists workspace_upstreams_workspace_base_auth_idx on workspace_upstreams(workspace_id, base_url, coalesce(auth_profile_id, ''))",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query(
|
||||
"insert into workspace_upstreams (
|
||||
@@ -411,7 +466,7 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
and wu.name = 'Frankfurter'
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -427,13 +482,13 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
updated_at timestamptz not null
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
"create unique index if not exists secrets_workspace_name_idx on secrets(workspace_id, name)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -447,7 +502,7 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
primary key (secret_id, version)
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -464,7 +519,7 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
finished_at timestamptz null
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -483,7 +538,7 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
finished_at timestamptz null
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -501,13 +556,13 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
published_at timestamptz null
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
"create unique index if not exists agents_workspace_slug_idx on agents(workspace_id, slug)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -521,7 +576,7 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
primary key (agent_id, version)
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -538,13 +593,13 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
foreign key (operation_id, operation_version) references operation_versions(operation_id, version) on delete cascade
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
"create unique index if not exists agent_bindings_tool_name_idx on agent_operation_bindings(agent_id, agent_version, tool_name)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -556,7 +611,7 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
foreign key (agent_id, version) references agent_versions(agent_id, version) on delete cascade
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -577,36 +632,30 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
decision_note text null
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query("alter table approval_requests drop column if exists confirmation_title")
|
||||
.execute(pool)
|
||||
.await?;
|
||||
query("alter table approval_requests drop column if exists confirmation_body")
|
||||
.execute(pool)
|
||||
.await?;
|
||||
query("alter table approval_requests add column if not exists execution_started_at timestamptz null")
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query("alter table approval_requests add column if not exists execution_attempts integer not null default 0")
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query("alter table approval_requests add column if not exists request_fingerprint text null")
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query(
|
||||
"create unique index if not exists approval_requests_pending_fingerprint_idx
|
||||
on approval_requests(agent_id, operation_id, operation_version, request_fingerprint)
|
||||
where status = 'pending' and request_fingerprint is not null",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
query(
|
||||
"create index if not exists approval_requests_agent_status_idx
|
||||
on approval_requests(workspace_id, agent_id, status, expires_at)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -629,25 +678,25 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
created_at timestamptz not null
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
"create index if not exists invocation_logs_workspace_created_idx on invocation_logs(workspace_id, created_at desc)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
"create index if not exists invocation_logs_workspace_operation_created_idx on invocation_logs(workspace_id, operation_id, created_at desc)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
"create index if not exists invocation_logs_workspace_agent_created_idx on invocation_logs(workspace_id, agent_id, created_at desc)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
query(
|
||||
@@ -665,7 +714,7 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> {
|
||||
updated_at timestamptz not null
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.execute(&mut **transaction)
|
||||
.await?;
|
||||
|
||||
Ok(())
|
||||
|
||||
@@ -395,11 +395,92 @@ pub struct FinishImportJobRequest<'a> {
|
||||
pub finished_at: &'a OffsetDateTime,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum ImportConflictMode {
|
||||
Skip,
|
||||
Rename,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ImportOperationDraft {
|
||||
pub operation_key: String,
|
||||
pub operation: RegistryOperation,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct AppliedImportOperation {
|
||||
pub operation_key: String,
|
||||
pub operation_id: OperationId,
|
||||
pub name: String,
|
||||
pub version: u32,
|
||||
pub renamed_from: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct SkippedImportOperation {
|
||||
pub operation_key: String,
|
||||
pub name: String,
|
||||
pub reason: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct ImportJobApplyResult {
|
||||
pub application_key: String,
|
||||
pub created: Vec<AppliedImportOperation>,
|
||||
pub skipped: Vec<SkippedImportOperation>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub struct ApplyImportJobRequest<'a> {
|
||||
pub id: &'a ImportJobId,
|
||||
pub workspace_id: &'a WorkspaceId,
|
||||
pub application_key: &'a str,
|
||||
pub conflict_mode: ImportConflictMode,
|
||||
pub operations: &'a [ImportOperationDraft],
|
||||
pub finished_at: &'a OffsetDateTime,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub struct CreateInvocationLogRequest<'a> {
|
||||
pub log: &'a InvocationLog,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub enum InvocationHistoryLossCategory {
|
||||
Unavailable,
|
||||
InvalidRecord,
|
||||
}
|
||||
|
||||
impl InvocationHistoryLossCategory {
|
||||
pub fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Unavailable => "unavailable",
|
||||
Self::InvalidRecord => "invalid_record",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub struct InvocationHistoryLoss {
|
||||
pub category: InvocationHistoryLossCategory,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub enum InvocationHistoryWriteOutcome {
|
||||
Recorded,
|
||||
Lost(InvocationHistoryLoss),
|
||||
}
|
||||
|
||||
impl InvocationHistoryWriteOutcome {
|
||||
pub fn loss(self) -> Option<InvocationHistoryLoss> {
|
||||
match self {
|
||||
Self::Recorded => None,
|
||||
Self::Lost(loss) => Some(loss),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct ListInvocationLogsQuery<'a> {
|
||||
pub workspace_id: &'a WorkspaceId,
|
||||
|
||||
@@ -397,18 +397,31 @@ impl PostgresRegistry {
|
||||
key_id: &PlatformApiKeyId,
|
||||
used_at: &time::OffsetDateTime,
|
||||
) -> Result<(), RegistryError> {
|
||||
let result = sqlx::query(
|
||||
"update platform_api_keys
|
||||
set last_used_at = $3::timestamptz
|
||||
where workspace_id = $1 and id = $2",
|
||||
let exists = sqlx::query_scalar::<_, bool>(
|
||||
"with target as (
|
||||
select id
|
||||
from platform_api_keys
|
||||
where workspace_id = $1 and id = $2
|
||||
), updated as (
|
||||
update platform_api_keys
|
||||
set last_used_at = $3::timestamptz
|
||||
where workspace_id = $1
|
||||
and id = $2
|
||||
and (
|
||||
last_used_at is null
|
||||
or last_used_at < $3::timestamptz - interval '1 minute'
|
||||
)
|
||||
returning id
|
||||
)
|
||||
select exists(select 1 from target)",
|
||||
)
|
||||
.bind(workspace_id.as_str())
|
||||
.bind(key_id.as_str())
|
||||
.bind(used_at)
|
||||
.execute(&self.pool)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
|
||||
if result.rows_affected() == 0 {
|
||||
if !exists {
|
||||
return Err(RegistryError::PlatformApiKeyNotFound {
|
||||
key_id: key_id.as_str().to_owned(),
|
||||
});
|
||||
|
||||
@@ -357,21 +357,20 @@ impl PostgresRegistry {
|
||||
&self,
|
||||
started_at: OffsetDateTime,
|
||||
approved_before: OffsetDateTime,
|
||||
stale_before: OffsetDateTime,
|
||||
) -> Result<Option<ApprovalRequestRecord>, RegistryError> {
|
||||
let row = sqlx::query(
|
||||
"with candidate as (
|
||||
select id
|
||||
from approval_requests
|
||||
where (status = 'approved' and decided_at <= $1)
|
||||
or (status = 'executing' and execution_started_at < $2)
|
||||
where status = 'approved'
|
||||
and decided_at <= $1
|
||||
order by decided_at asc nulls last, created_at asc
|
||||
for update skip locked
|
||||
limit 1
|
||||
)
|
||||
update approval_requests as approval
|
||||
set status = 'executing',
|
||||
execution_started_at = $3,
|
||||
execution_started_at = $2,
|
||||
execution_attempts = approval.execution_attempts + 1
|
||||
from candidate
|
||||
where approval.id = candidate.id
|
||||
@@ -383,7 +382,6 @@ impl PostgresRegistry {
|
||||
approval.decided_at, approval.decided_by_key_id, approval.decision_note",
|
||||
)
|
||||
.bind(approved_before)
|
||||
.bind(stale_before)
|
||||
.bind(started_at)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
@@ -391,6 +389,50 @@ impl PostgresRegistry {
|
||||
row.map(map_approval_request_row).transpose()
|
||||
}
|
||||
|
||||
pub async fn fail_next_interrupted_approval_request(
|
||||
&self,
|
||||
stale_before: OffsetDateTime,
|
||||
) -> Result<Option<ApprovalRequestRecord>, RegistryError> {
|
||||
let response_payload = serde_json::json!({
|
||||
"error": {
|
||||
"code": "approval_execution_outcome_unknown",
|
||||
"message": "execution was interrupted; the operation was not retried automatically"
|
||||
}
|
||||
});
|
||||
let row = sqlx::query(
|
||||
"with candidate as (
|
||||
select id
|
||||
from approval_requests
|
||||
where status = 'executing'
|
||||
and execution_started_at < $1
|
||||
order by execution_started_at asc, created_at asc
|
||||
for update skip locked
|
||||
limit 1
|
||||
)
|
||||
update approval_requests as approval
|
||||
set status = 'failed',
|
||||
response_payload_json = $2,
|
||||
decision_note = coalesce(
|
||||
approval.decision_note,
|
||||
'execution interrupted; outcome unknown'
|
||||
)
|
||||
from candidate
|
||||
where approval.id = candidate.id
|
||||
returning
|
||||
approval.id, approval.workspace_id, approval.agent_id,
|
||||
approval.operation_id, approval.operation_version, approval.status,
|
||||
approval.risk_level, approval.request_payload_json,
|
||||
approval.response_payload_json, approval.created_at, approval.expires_at,
|
||||
approval.decided_at, approval.decided_by_key_id, approval.decision_note",
|
||||
)
|
||||
.bind(stale_before)
|
||||
.bind(Json(response_payload))
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
|
||||
row.map(map_approval_request_row).transpose()
|
||||
}
|
||||
|
||||
pub async fn expire_approval_request(
|
||||
&self,
|
||||
request: ExpireApprovalRequest<'_>,
|
||||
|
||||
@@ -1,15 +1,31 @@
|
||||
use super::*;
|
||||
use time::OffsetDateTime;
|
||||
|
||||
fn map_auth_user_row(row: &PgRow) -> Result<AuthUserRecord, RegistryError> {
|
||||
let status = row.try_get::<String, _>("status")?;
|
||||
Ok(AuthUserRecord {
|
||||
user: User {
|
||||
id: UserId::new(row.try_get::<String, _>("id")?),
|
||||
email: row.try_get("email")?,
|
||||
display_name: row.try_get("display_name")?,
|
||||
status: deserialize_enum_text(&status, "status")?,
|
||||
created_at: row.try_get::<OffsetDateTime, _>("created_at")?,
|
||||
},
|
||||
password_hash: row
|
||||
.try_get::<Option<String>, _>("password_hash")?
|
||||
.unwrap_or_default(),
|
||||
})
|
||||
}
|
||||
|
||||
impl PostgresRegistry {
|
||||
pub async fn upsert_bootstrap_user(
|
||||
pub async fn ensure_bootstrap_user(
|
||||
&self,
|
||||
email: &str,
|
||||
display_name: &str,
|
||||
password_hash: &str,
|
||||
) -> Result<UserId, RegistryError> {
|
||||
let user_id = format!("user_{}", uuid::Uuid::now_v7().simple());
|
||||
let row = sqlx::query!(
|
||||
if let Some(id) = sqlx::query_scalar::<_, String>(
|
||||
"insert into users (
|
||||
id,
|
||||
email,
|
||||
@@ -24,16 +40,40 @@ impl PostgresRegistry {
|
||||
set display_name = excluded.display_name,
|
||||
password_hash = excluded.password_hash,
|
||||
status = 'active'
|
||||
where users.password_hash is null
|
||||
returning id",
|
||||
user_id,
|
||||
email,
|
||||
display_name,
|
||||
password_hash,
|
||||
)
|
||||
.bind(user_id)
|
||||
.bind(email)
|
||||
.bind(display_name)
|
||||
.bind(password_hash)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?
|
||||
{
|
||||
return Ok(UserId::new(id));
|
||||
}
|
||||
|
||||
let existing = sqlx::query_scalar::<_, String>(
|
||||
"select id
|
||||
from users
|
||||
where email = $1
|
||||
limit 1",
|
||||
)
|
||||
.bind(email)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
|
||||
Ok(UserId::new(row.id))
|
||||
Ok(UserId::new(existing))
|
||||
}
|
||||
|
||||
pub async fn upsert_bootstrap_user(
|
||||
&self,
|
||||
email: &str,
|
||||
display_name: &str,
|
||||
password_hash: &str,
|
||||
) -> Result<UserId, RegistryError> {
|
||||
self.ensure_bootstrap_user(email, display_name, password_hash)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn ensure_membership(
|
||||
@@ -67,70 +107,46 @@ impl PostgresRegistry {
|
||||
&self,
|
||||
email: &str,
|
||||
) -> Result<Option<AuthUserRecord>, RegistryError> {
|
||||
let row = sqlx::query!(
|
||||
let row = sqlx::query(
|
||||
"select
|
||||
id,
|
||||
email,
|
||||
display_name,
|
||||
password_hash as \"password_hash!\",
|
||||
password_hash,
|
||||
status,
|
||||
created_at as \"created_at!: OffsetDateTime\"
|
||||
created_at
|
||||
from users
|
||||
where email = $1
|
||||
limit 1",
|
||||
email,
|
||||
)
|
||||
.bind(email)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
|
||||
row.map(|row| {
|
||||
Ok(AuthUserRecord {
|
||||
user: User {
|
||||
id: UserId::new(row.id),
|
||||
email: row.email,
|
||||
display_name: row.display_name,
|
||||
status: deserialize_enum_text(&row.status, "status")?,
|
||||
created_at: row.created_at,
|
||||
},
|
||||
password_hash: row.password_hash,
|
||||
})
|
||||
})
|
||||
.transpose()
|
||||
row.as_ref().map(map_auth_user_row).transpose()
|
||||
}
|
||||
|
||||
pub async fn get_auth_user_by_id(
|
||||
&self,
|
||||
user_id: &UserId,
|
||||
) -> Result<Option<AuthUserRecord>, RegistryError> {
|
||||
let row = sqlx::query!(
|
||||
let row = sqlx::query(
|
||||
"select
|
||||
id,
|
||||
email,
|
||||
display_name,
|
||||
password_hash as \"password_hash!\",
|
||||
password_hash,
|
||||
status,
|
||||
created_at as \"created_at!: OffsetDateTime\"
|
||||
created_at
|
||||
from users
|
||||
where id = $1
|
||||
limit 1",
|
||||
user_id.as_str(),
|
||||
)
|
||||
.bind(user_id.as_str())
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
|
||||
row.map(|row| {
|
||||
Ok(AuthUserRecord {
|
||||
user: User {
|
||||
id: UserId::new(row.id),
|
||||
email: row.email,
|
||||
display_name: row.display_name,
|
||||
status: deserialize_enum_text(&row.status, "status")?,
|
||||
created_at: row.created_at,
|
||||
},
|
||||
password_hash: row.password_hash,
|
||||
})
|
||||
})
|
||||
.transpose()
|
||||
row.as_ref().map(map_auth_user_row).transpose()
|
||||
}
|
||||
|
||||
pub async fn update_user_profile(
|
||||
@@ -191,6 +207,45 @@ impl PostgresRegistry {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn update_user_password_and_revoke_other_sessions(
|
||||
&self,
|
||||
user_id: &UserId,
|
||||
current_session_id: &UserSessionId,
|
||||
password_hash: &str,
|
||||
) -> Result<(), RegistryError> {
|
||||
let mut tx = self.pool.begin().await?;
|
||||
let result = sqlx::query(
|
||||
"update users
|
||||
set password_hash = $2
|
||||
where id = $1",
|
||||
)
|
||||
.bind(user_id.as_str())
|
||||
.bind(password_hash)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
|
||||
if result.rows_affected() == 0 {
|
||||
return Err(RegistryError::UserNotFound {
|
||||
user_id: user_id.as_str().to_owned(),
|
||||
});
|
||||
}
|
||||
|
||||
sqlx::query(
|
||||
"update user_sessions
|
||||
set status = 'revoked'
|
||||
where user_id = $1
|
||||
and id <> $2
|
||||
and status = 'active'",
|
||||
)
|
||||
.bind(user_id.as_str())
|
||||
.bind(current_session_id.as_str())
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
|
||||
tx.commit().await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn create_user_session(
|
||||
&self,
|
||||
session_id: &UserSessionId,
|
||||
@@ -305,7 +360,11 @@ impl PostgresRegistry {
|
||||
sqlx::query(
|
||||
"update user_sessions
|
||||
set last_seen_at = now()
|
||||
where id = $1",
|
||||
where id = $1
|
||||
and (
|
||||
last_seen_at is null
|
||||
or last_seen_at < now() - interval '1 minute'
|
||||
)",
|
||||
)
|
||||
.bind(session_id.as_str())
|
||||
.execute(&self.pool)
|
||||
|
||||
@@ -41,6 +41,11 @@ impl PostgresRegistry {
|
||||
&self.pool
|
||||
}
|
||||
|
||||
pub async fn ping(&self) -> Result<(), RegistryError> {
|
||||
sqlx::query("select 1").execute(&self.pool).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn migrate(&self) -> Result<(), RegistryError> {
|
||||
migrations::apply_postgres(&self.pool).await?;
|
||||
Ok(())
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
use super::*;
|
||||
|
||||
const APPLICATION_RESULT_KEY: &str = "_crank_application_result";
|
||||
|
||||
impl PostgresRegistry {
|
||||
pub async fn create_import_job(
|
||||
&self,
|
||||
@@ -97,6 +99,43 @@ impl PostgresRegistry {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn apply_import_job(
|
||||
&self,
|
||||
request: ApplyImportJobRequest<'_>,
|
||||
) -> Result<ImportJobApplyResult, RegistryError> {
|
||||
let mut tx = self.pool.begin().await?;
|
||||
let applied = apply_import_job_transaction(&mut tx, &request).await;
|
||||
|
||||
match applied {
|
||||
Ok(result) => {
|
||||
tx.commit().await?;
|
||||
Ok(result)
|
||||
}
|
||||
Err(error) => {
|
||||
tx.rollback().await?;
|
||||
let error_text = error.to_string();
|
||||
let _ = sqlx::query(
|
||||
"update import_jobs
|
||||
set status = $3,
|
||||
error_text = $4,
|
||||
finished_at = $5::timestamptz
|
||||
where id = $1
|
||||
and workspace_id = $2
|
||||
and status <> $6",
|
||||
)
|
||||
.bind(request.id.as_str())
|
||||
.bind(request.workspace_id.as_str())
|
||||
.bind(serialize_enum_text(&ImportJobStatus::Failed, "status")?)
|
||||
.bind(error_text)
|
||||
.bind(request.finished_at)
|
||||
.bind(serialize_enum_text(&ImportJobStatus::Completed, "status")?)
|
||||
.execute(&self.pool)
|
||||
.await;
|
||||
Err(error)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn delete_expired_import_jobs(&self) -> Result<u64, RegistryError> {
|
||||
let result = sqlx::query("delete from import_jobs where expires_at < now()")
|
||||
.execute(&self.pool)
|
||||
@@ -105,3 +144,160 @@ impl PostgresRegistry {
|
||||
Ok(result.rows_affected())
|
||||
}
|
||||
}
|
||||
|
||||
async fn apply_import_job_transaction(
|
||||
tx: &mut Transaction<'_, Postgres>,
|
||||
request: &ApplyImportJobRequest<'_>,
|
||||
) -> Result<ImportJobApplyResult, RegistryError> {
|
||||
let row = sqlx::query(
|
||||
"select status, preview_payload
|
||||
from import_jobs
|
||||
where id = $1 and workspace_id = $2
|
||||
for update",
|
||||
)
|
||||
.bind(request.id.as_str())
|
||||
.bind(request.workspace_id.as_str())
|
||||
.fetch_optional(&mut **tx)
|
||||
.await?
|
||||
.ok_or_else(|| RegistryError::ImportJobNotFound {
|
||||
job_id: request.id.as_str().to_owned(),
|
||||
})?;
|
||||
let status = deserialize_enum_text::<ImportJobStatus>(row.try_get("status")?, "status")?;
|
||||
let mut preview_payload = row.try_get::<Value, _>("preview_payload")?;
|
||||
|
||||
if status == ImportJobStatus::Completed
|
||||
&& let Some(result) = stored_application_result(&preview_payload)?
|
||||
{
|
||||
if result.application_key != request.application_key {
|
||||
return Err(RegistryError::ImportJobAlreadyApplied {
|
||||
job_id: request.id.as_str().to_owned(),
|
||||
});
|
||||
}
|
||||
return Ok(result);
|
||||
}
|
||||
|
||||
sqlx::query("select id from workspaces where id = $1 for update")
|
||||
.bind(request.workspace_id.as_str())
|
||||
.fetch_one(&mut **tx)
|
||||
.await?;
|
||||
|
||||
let mut result = ImportJobApplyResult {
|
||||
application_key: request.application_key.to_owned(),
|
||||
..ImportJobApplyResult::default()
|
||||
};
|
||||
for draft in request.operations {
|
||||
if draft.operation.version != 1 {
|
||||
return Err(RegistryError::InvalidInitialVersion {
|
||||
operation_id: draft.operation.id.as_str().to_owned(),
|
||||
version: draft.operation.version,
|
||||
});
|
||||
}
|
||||
|
||||
let mut operation = draft.operation.clone();
|
||||
let original_name = operation.name.clone();
|
||||
if operation_name_exists(tx, request.workspace_id, &operation.name).await? {
|
||||
match request.conflict_mode {
|
||||
ImportConflictMode::Skip => {
|
||||
result.skipped.push(SkippedImportOperation {
|
||||
operation_key: draft.operation_key.clone(),
|
||||
name: operation.name,
|
||||
reason: "operation_name_conflict".to_owned(),
|
||||
});
|
||||
continue;
|
||||
}
|
||||
ImportConflictMode::Rename => {
|
||||
operation.name =
|
||||
next_available_operation_name(tx, request.workspace_id, &operation.name)
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
insert_operation_rows(tx, request.workspace_id, &operation, None).await?;
|
||||
result.created.push(AppliedImportOperation {
|
||||
operation_key: draft.operation_key.clone(),
|
||||
operation_id: operation.id,
|
||||
name: operation.name.clone(),
|
||||
version: operation.version,
|
||||
renamed_from: (operation.name != original_name).then_some(original_name),
|
||||
});
|
||||
}
|
||||
|
||||
let stored_result = serde_json::to_value(&result)?;
|
||||
if let Some(object) = preview_payload.as_object_mut() {
|
||||
object.insert(APPLICATION_RESULT_KEY.to_owned(), stored_result);
|
||||
} else {
|
||||
preview_payload = serde_json::json!({
|
||||
"preview": preview_payload,
|
||||
"_crank_application_result": stored_result,
|
||||
});
|
||||
}
|
||||
let created_operation_ids = serde_json::to_value(
|
||||
result
|
||||
.created
|
||||
.iter()
|
||||
.map(|operation| operation.operation_id.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
)?;
|
||||
sqlx::query(
|
||||
"update import_jobs
|
||||
set status = $3,
|
||||
preview_payload = $4,
|
||||
created_operation_ids = $5,
|
||||
error_text = null,
|
||||
finished_at = $6::timestamptz
|
||||
where id = $1 and workspace_id = $2",
|
||||
)
|
||||
.bind(request.id.as_str())
|
||||
.bind(request.workspace_id.as_str())
|
||||
.bind(serialize_enum_text(&ImportJobStatus::Completed, "status")?)
|
||||
.bind(preview_payload)
|
||||
.bind(created_operation_ids)
|
||||
.bind(request.finished_at)
|
||||
.execute(&mut **tx)
|
||||
.await?;
|
||||
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
fn stored_application_result(
|
||||
preview_payload: &Value,
|
||||
) -> Result<Option<ImportJobApplyResult>, RegistryError> {
|
||||
preview_payload
|
||||
.get(APPLICATION_RESULT_KEY)
|
||||
.cloned()
|
||||
.map(serde_json::from_value)
|
||||
.transpose()
|
||||
.map_err(RegistryError::from)
|
||||
}
|
||||
|
||||
async fn operation_name_exists(
|
||||
tx: &mut Transaction<'_, Postgres>,
|
||||
workspace_id: &WorkspaceId,
|
||||
name: &str,
|
||||
) -> Result<bool, RegistryError> {
|
||||
Ok(sqlx::query_scalar::<_, bool>(
|
||||
"select exists(
|
||||
select 1 from operations where workspace_id = $1 and name = $2
|
||||
)",
|
||||
)
|
||||
.bind(workspace_id.as_str())
|
||||
.bind(name)
|
||||
.fetch_one(&mut **tx)
|
||||
.await?)
|
||||
}
|
||||
|
||||
async fn next_available_operation_name(
|
||||
tx: &mut Transaction<'_, Postgres>,
|
||||
workspace_id: &WorkspaceId,
|
||||
base_name: &str,
|
||||
) -> Result<String, RegistryError> {
|
||||
for index in 2.. {
|
||||
let candidate = format!("{base_name}_{index}");
|
||||
if !operation_name_exists(tx, workspace_id, &candidate).await? {
|
||||
return Ok(candidate);
|
||||
}
|
||||
}
|
||||
|
||||
unreachable!()
|
||||
}
|
||||
|
||||
@@ -30,23 +30,25 @@ pub use pool_config::{PostgresPoolConfig, PostgresPoolConfigError};
|
||||
use crate::{
|
||||
error::RegistryError,
|
||||
model::{
|
||||
AgentSummary, AgentVersionRecord, ApprovalRequestRecord, AuthUserRecord,
|
||||
CreateAgentDraftVersionRequest, CreateAgentRequest, CreateApprovalRequest,
|
||||
CreateImportJobRequest, CreateInvitationRequest, CreateInvocationLogRequest,
|
||||
CreatePlatformApiKeyRequest, CreateSecretRequest, CreateVersionRequest,
|
||||
CreateWorkspaceRequest, CreateYamlImportJobRequest, DecideApprovalRequest,
|
||||
DescriptorMetadata, ExpireApprovalRequest, FinishApprovalRequest, FinishImportJobRequest,
|
||||
ImportJob, ImportJobId, InvitationRecord, InvocationLogRecord, ListApprovalRequestsQuery,
|
||||
AgentSummary, AgentVersionRecord, AppliedImportOperation, ApplyImportJobRequest,
|
||||
ApprovalRequestRecord, AuthUserRecord, CreateAgentDraftVersionRequest, CreateAgentRequest,
|
||||
CreateApprovalRequest, CreateImportJobRequest, CreateInvitationRequest,
|
||||
CreateInvocationLogRequest, CreatePlatformApiKeyRequest, CreateSecretRequest,
|
||||
CreateVersionRequest, CreateWorkspaceRequest, CreateYamlImportJobRequest,
|
||||
DecideApprovalRequest, DescriptorMetadata, ExpireApprovalRequest, FinishApprovalRequest,
|
||||
FinishImportJobRequest, ImportConflictMode, ImportJob, ImportJobApplyResult, ImportJobId,
|
||||
ImportJobStatus, InvitationRecord, InvocationHistoryLoss, InvocationHistoryLossCategory,
|
||||
InvocationHistoryWriteOutcome, InvocationLogRecord, ListApprovalRequestsQuery,
|
||||
ListInvocationLogsQuery, MembershipRecord, OperationAgentRef, OperationSampleMetadata,
|
||||
OperationSummary, OperationUsageSummary, OperationVersionRecord, PlatformApiKeyRecord,
|
||||
PublishAgentRequest, PublishRequest, PublishedAgentCatalog, PublishedAgentTool,
|
||||
RegistryOperation, RotateSecretRequest, SaveAgentBindingsRequest,
|
||||
SaveAgentCatalogConfigRequest, SaveAuthProfileRequest, SaveDescriptorMetadataRequest,
|
||||
SaveSampleMetadataRequest, SaveWorkspaceUpstreamRequest, SecretRecord, SecretVersionRecord,
|
||||
SessionRecord, UpdateWorkspaceRequest, UsageAgentBreakdown, UsageOperationBreakdown,
|
||||
UsageQuery, UsageRollupRecord, UsageSummary, UsageTimelinePoint, WorkspaceMembershipRecord,
|
||||
WorkspaceRecord, WorkspaceUpstream, YamlImportJob, YamlImportJobCompletion,
|
||||
YamlImportJobId, YamlImportJobStatus,
|
||||
SessionRecord, SkippedImportOperation, UpdateWorkspaceRequest, UsageAgentBreakdown,
|
||||
UsageOperationBreakdown, UsageQuery, UsageRollupRecord, UsageSummary, UsageTimelinePoint,
|
||||
WorkspaceMembershipRecord, WorkspaceRecord, WorkspaceUpstream, YamlImportJob,
|
||||
YamlImportJobCompletion, YamlImportJobId, YamlImportJobStatus,
|
||||
},
|
||||
};
|
||||
|
||||
@@ -135,6 +137,61 @@ async fn insert_version_row(
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn insert_operation_rows(
|
||||
tx: &mut Transaction<'_, Postgres>,
|
||||
workspace_id: &WorkspaceId,
|
||||
snapshot: &RegistryOperation,
|
||||
created_by: Option<&str>,
|
||||
) -> Result<(), RegistryError> {
|
||||
sqlx::query(
|
||||
"insert into operations (
|
||||
id,
|
||||
workspace_id,
|
||||
name,
|
||||
display_name,
|
||||
category,
|
||||
protocol,
|
||||
security_level,
|
||||
status,
|
||||
current_draft_version,
|
||||
latest_published_version,
|
||||
created_at,
|
||||
updated_at,
|
||||
published_at
|
||||
) values (
|
||||
$1, $2, $3, $4, $5, $6, $7, $8, $9, $10,
|
||||
$11::timestamptz,
|
||||
$12::timestamptz,
|
||||
$13::timestamptz
|
||||
)",
|
||||
)
|
||||
.bind(snapshot.id.as_str())
|
||||
.bind(workspace_id.as_str())
|
||||
.bind(&snapshot.name)
|
||||
.bind(&snapshot.display_name)
|
||||
.bind(&snapshot.category)
|
||||
.bind(serialize_enum_text(&snapshot.protocol, "protocol")?)
|
||||
.bind(serialize_enum_text(
|
||||
&snapshot.security_level,
|
||||
"security_level",
|
||||
)?)
|
||||
.bind(serialize_enum_text(&snapshot.status, "status")?)
|
||||
.bind(to_db_version(snapshot.version))
|
||||
.bind(
|
||||
snapshot
|
||||
.published_at
|
||||
.as_ref()
|
||||
.map(|_| to_db_version(snapshot.version)),
|
||||
)
|
||||
.bind(snapshot.created_at)
|
||||
.bind(snapshot.updated_at)
|
||||
.bind(snapshot.published_at)
|
||||
.execute(&mut **tx)
|
||||
.await?;
|
||||
|
||||
insert_version_row(tx, snapshot, None, created_by).await
|
||||
}
|
||||
|
||||
async fn insert_agent_version_row(
|
||||
tx: &mut Transaction<'_, Postgres>,
|
||||
version: &AgentVersion,
|
||||
|
||||
@@ -1,9 +1,47 @@
|
||||
use super::*;
|
||||
|
||||
fn invocation_history_loss_category(error: &RegistryError) -> InvocationHistoryLossCategory {
|
||||
match error {
|
||||
RegistryError::Storage(error)
|
||||
if error
|
||||
.as_database_error()
|
||||
.and_then(|database_error| database_error.code())
|
||||
.is_some_and(|code| code.starts_with("22") || code.starts_with("23")) =>
|
||||
{
|
||||
InvocationHistoryLossCategory::InvalidRecord
|
||||
}
|
||||
RegistryError::Storage(_) => InvocationHistoryLossCategory::Unavailable,
|
||||
_ => InvocationHistoryLossCategory::InvalidRecord,
|
||||
}
|
||||
}
|
||||
|
||||
impl PostgresRegistry {
|
||||
pub async fn delete_invocation_logs_before(
|
||||
&self,
|
||||
cutoff: OffsetDateTime,
|
||||
) -> Result<u64, RegistryError> {
|
||||
let result = sqlx::query("delete from invocation_logs where created_at < $1::timestamptz")
|
||||
.bind(cutoff)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
Ok(result.rows_affected())
|
||||
}
|
||||
|
||||
pub async fn create_invocation_log(
|
||||
&self,
|
||||
request: CreateInvocationLogRequest<'_>,
|
||||
) -> InvocationHistoryWriteOutcome {
|
||||
match self.try_create_invocation_log(request).await {
|
||||
Ok(()) => InvocationHistoryWriteOutcome::Recorded,
|
||||
Err(error) => InvocationHistoryWriteOutcome::Lost(InvocationHistoryLoss {
|
||||
category: invocation_history_loss_category(&error),
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
async fn try_create_invocation_log(
|
||||
&self,
|
||||
request: CreateInvocationLogRequest<'_>,
|
||||
) -> Result<(), RegistryError> {
|
||||
let request_preview = crank_core::sanitize_invocation_preview(&request.log.request_preview);
|
||||
let response_preview =
|
||||
|
||||
@@ -26,53 +26,7 @@ impl PostgresRegistry {
|
||||
|
||||
let mut tx = self.pool.begin().await?;
|
||||
|
||||
sqlx::query(
|
||||
"insert into operations (
|
||||
id,
|
||||
workspace_id,
|
||||
name,
|
||||
display_name,
|
||||
category,
|
||||
protocol,
|
||||
security_level,
|
||||
status,
|
||||
current_draft_version,
|
||||
latest_published_version,
|
||||
created_at,
|
||||
updated_at,
|
||||
published_at
|
||||
) values (
|
||||
$1, $2, $3, $4, $5, $6, $7, $8, $9, $10,
|
||||
$11::timestamptz,
|
||||
$12::timestamptz,
|
||||
$13::timestamptz
|
||||
)",
|
||||
)
|
||||
.bind(snapshot.id.as_str())
|
||||
.bind(workspace_id.as_str())
|
||||
.bind(&snapshot.name)
|
||||
.bind(&snapshot.display_name)
|
||||
.bind(&snapshot.category)
|
||||
.bind(serialize_enum_text(&snapshot.protocol, "protocol")?)
|
||||
.bind(serialize_enum_text(
|
||||
&snapshot.security_level,
|
||||
"security_level",
|
||||
)?)
|
||||
.bind(serialize_enum_text(&snapshot.status, "status")?)
|
||||
.bind(to_db_version(snapshot.version))
|
||||
.bind(
|
||||
snapshot
|
||||
.published_at
|
||||
.as_ref()
|
||||
.map(|_| to_db_version(snapshot.version)),
|
||||
)
|
||||
.bind(snapshot.created_at)
|
||||
.bind(snapshot.updated_at)
|
||||
.bind(snapshot.published_at)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
|
||||
insert_version_row(&mut tx, snapshot, None, created_by).await?;
|
||||
insert_operation_rows(&mut tx, workspace_id, snapshot, created_by).await?;
|
||||
|
||||
tx.commit().await?;
|
||||
Ok(())
|
||||
|
||||
@@ -282,18 +282,31 @@ impl PostgresRegistry {
|
||||
secret_id: &SecretId,
|
||||
used_at: &OffsetDateTime,
|
||||
) -> Result<(), RegistryError> {
|
||||
let result = sqlx::query(
|
||||
"update secrets
|
||||
set last_used_at = $3::timestamptz
|
||||
where workspace_id = $1 and id = $2",
|
||||
let exists = sqlx::query_scalar::<_, bool>(
|
||||
"with target as (
|
||||
select id
|
||||
from secrets
|
||||
where workspace_id = $1 and id = $2
|
||||
), updated as (
|
||||
update secrets
|
||||
set last_used_at = $3::timestamptz
|
||||
where workspace_id = $1
|
||||
and id = $2
|
||||
and (
|
||||
last_used_at is null
|
||||
or last_used_at < $3::timestamptz - interval '1 minute'
|
||||
)
|
||||
returning id
|
||||
)
|
||||
select exists(select 1 from target)",
|
||||
)
|
||||
.bind(workspace_id.as_str())
|
||||
.bind(secret_id.as_str())
|
||||
.bind(used_at)
|
||||
.execute(&self.pool)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
|
||||
if result.rows_affected() == 0 {
|
||||
if !exists {
|
||||
return Err(RegistryError::SecretNotFound {
|
||||
secret_id: secret_id.as_str().to_owned(),
|
||||
});
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
mod integration {
|
||||
mod agents_usage;
|
||||
mod common;
|
||||
mod migrations;
|
||||
mod observability;
|
||||
mod operations_artifacts;
|
||||
mod workspace_access;
|
||||
}
|
||||
|
||||
@@ -228,32 +228,36 @@ async fn manages_operation_usage_and_agent_ref_reads() {
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
registry
|
||||
.create_invocation_log(CreateInvocationLogRequest {
|
||||
log: &test_invocation_log(
|
||||
"log_usage_ok",
|
||||
&operation.id,
|
||||
Some(agent.id.clone()),
|
||||
crank_core::InvocationStatus::Ok,
|
||||
120,
|
||||
"2026-03-25T12:20:00Z",
|
||||
),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
registry
|
||||
.create_invocation_log(CreateInvocationLogRequest {
|
||||
log: &test_invocation_log(
|
||||
"log_usage_err",
|
||||
&operation.id,
|
||||
Some(agent.id.clone()),
|
||||
crank_core::InvocationStatus::Error,
|
||||
240,
|
||||
"2026-03-25T12:21:00Z",
|
||||
),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
registry
|
||||
.create_invocation_log(CreateInvocationLogRequest {
|
||||
log: &test_invocation_log(
|
||||
"log_usage_ok",
|
||||
&operation.id,
|
||||
Some(agent.id.clone()),
|
||||
crank_core::InvocationStatus::Ok,
|
||||
120,
|
||||
"2026-03-25T12:20:00Z",
|
||||
),
|
||||
})
|
||||
.await,
|
||||
crank_registry::InvocationHistoryWriteOutcome::Recorded
|
||||
);
|
||||
assert_eq!(
|
||||
registry
|
||||
.create_invocation_log(CreateInvocationLogRequest {
|
||||
log: &test_invocation_log(
|
||||
"log_usage_err",
|
||||
&operation.id,
|
||||
Some(agent.id.clone()),
|
||||
crank_core::InvocationStatus::Error,
|
||||
240,
|
||||
"2026-03-25T12:21:00Z",
|
||||
),
|
||||
})
|
||||
.await,
|
||||
crank_registry::InvocationHistoryWriteOutcome::Recorded
|
||||
);
|
||||
|
||||
let has_bindings = registry
|
||||
.has_published_agent_bindings_for_operation(&test_workspace_id(), &operation.id)
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
use crank_registry::PostgresRegistry;
|
||||
use sqlx::Row;
|
||||
|
||||
#[tokio::test]
|
||||
async fn core_migration_is_versioned_and_safe_under_concurrent_startup() {
|
||||
let database_url = crank_test_support::postgres_schema_url("test_core_migration").await;
|
||||
|
||||
let (first, second) = tokio::join!(
|
||||
PostgresRegistry::connect(&database_url),
|
||||
PostgresRegistry::connect(&database_url),
|
||||
);
|
||||
let first = first.expect("first service startup must apply the migration");
|
||||
second.expect("second service startup must observe the applied migration");
|
||||
|
||||
let rows = sqlx::query(
|
||||
"select version, description, checksum from __crank_core_migrations order by version",
|
||||
)
|
||||
.fetch_all(first.pool())
|
||||
.await
|
||||
.expect("migration ledger must be readable");
|
||||
|
||||
assert_eq!(rows.len(), 1);
|
||||
assert_eq!(rows[0].get::<i32, _>("version"), 1);
|
||||
assert_eq!(
|
||||
rows[0].get::<String, _>("description"),
|
||||
"community baseline"
|
||||
);
|
||||
assert_eq!(
|
||||
rows[0].get::<String, _>("checksum"),
|
||||
"crank-community-baseline-v1"
|
||||
);
|
||||
|
||||
let approval_columns = sqlx::query(
|
||||
"select column_name
|
||||
from information_schema.columns
|
||||
where table_schema = current_schema()
|
||||
and table_name = 'approval_requests'
|
||||
and column_name in ('execution_started_at', 'execution_attempts', 'request_fingerprint')",
|
||||
)
|
||||
.fetch_all(first.pool())
|
||||
.await
|
||||
.expect("approval schema must be readable");
|
||||
assert_eq!(approval_columns.len(), 3);
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
use crank_registry::{
|
||||
CreateInvocationLogRequest, InvocationHistoryLossCategory, InvocationHistoryWriteOutcome,
|
||||
};
|
||||
|
||||
use super::common::{TestDatabase, test_invocation_log};
|
||||
|
||||
#[tokio::test]
|
||||
async fn invocation_history_write_returns_typed_loss_without_error_details() {
|
||||
let database = TestDatabase::new().await;
|
||||
let registry = database.registry().await;
|
||||
let log = test_invocation_log(
|
||||
"log_missing_owner",
|
||||
&crank_core::OperationId::new("op_missing"),
|
||||
None,
|
||||
crank_core::InvocationStatus::Ok,
|
||||
10,
|
||||
"2026-03-25T12:20:00Z",
|
||||
);
|
||||
|
||||
let outcome = registry
|
||||
.create_invocation_log(CreateInvocationLogRequest { log: &log })
|
||||
.await;
|
||||
|
||||
assert_eq!(
|
||||
outcome,
|
||||
InvocationHistoryWriteOutcome::Lost(crank_registry::InvocationHistoryLoss {
|
||||
category: InvocationHistoryLossCategory::InvalidRecord,
|
||||
})
|
||||
);
|
||||
database.cleanup().await;
|
||||
}
|
||||
@@ -38,6 +38,38 @@ fn timestamp(value: &str) -> OffsetDateTime {
|
||||
OffsetDateTime::parse(value, &Rfc3339).unwrap()
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn bootstrap_user_does_not_overwrite_an_existing_password() {
|
||||
let database = TestDatabase::new().await;
|
||||
let registry = database.registry().await;
|
||||
let email = "bootstrap-owner@example.com";
|
||||
|
||||
let user_id = registry
|
||||
.upsert_bootstrap_user(email, "Bootstrap Owner", "initial-bootstrap-hash")
|
||||
.await
|
||||
.unwrap();
|
||||
registry
|
||||
.update_user_password(&user_id, "user-selected-hash")
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let repeated_user_id = registry
|
||||
.upsert_bootstrap_user(email, "Changed Bootstrap Name", "changed-bootstrap-hash")
|
||||
.await
|
||||
.unwrap();
|
||||
let stored = registry
|
||||
.get_auth_user_by_email(email)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(repeated_user_id, user_id);
|
||||
assert_eq!(stored.password_hash, "user-selected-hash");
|
||||
assert_eq!(stored.user.display_name, "Bootstrap Owner");
|
||||
|
||||
database.cleanup().await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stores_and_finishes_yaml_import_jobs() {
|
||||
let database = TestDatabase::new().await;
|
||||
@@ -352,6 +384,24 @@ async fn creates_and_loads_user_sessions_with_typed_expiration() {
|
||||
assert_eq!(session.user.id, user_id);
|
||||
assert!(session.user.created_at.unix_timestamp() > 0);
|
||||
|
||||
registry.touch_user_session(&session_id).await.unwrap();
|
||||
let first_seen = sqlx::query_scalar::<_, OffsetDateTime>(
|
||||
"select last_seen_at from user_sessions where id = $1",
|
||||
)
|
||||
.bind(session_id.as_str())
|
||||
.fetch_one(registry.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
registry.touch_user_session(&session_id).await.unwrap();
|
||||
let second_seen = sqlx::query_scalar::<_, OffsetDateTime>(
|
||||
"select last_seen_at from user_sessions where id = $1",
|
||||
)
|
||||
.bind(session_id.as_str())
|
||||
.fetch_one(registry.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(second_seen, first_seen);
|
||||
|
||||
database.cleanup().await;
|
||||
}
|
||||
|
||||
@@ -454,6 +504,40 @@ async fn manages_platform_api_key_read_paths() {
|
||||
Some(timestamp("2026-03-25T12:05:00Z"))
|
||||
);
|
||||
|
||||
registry
|
||||
.touch_platform_api_key(
|
||||
&workspace.id,
|
||||
&PlatformApiKeyId::new("key_01"),
|
||||
×tamp("2026-03-25T12:05:30Z"),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let throttled = registry
|
||||
.list_platform_api_keys(&workspace.id)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
throttled[0].api_key.last_used_at,
|
||||
Some(timestamp("2026-03-25T12:05:00Z"))
|
||||
);
|
||||
|
||||
registry
|
||||
.touch_platform_api_key(
|
||||
&workspace.id,
|
||||
&PlatformApiKeyId::new("key_01"),
|
||||
×tamp("2026-03-25T12:06:01Z"),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let refreshed = registry
|
||||
.list_platform_api_keys(&workspace.id)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
refreshed[0].api_key.last_used_at,
|
||||
Some(timestamp("2026-03-25T12:06:01Z"))
|
||||
);
|
||||
|
||||
database.cleanup().await;
|
||||
}
|
||||
|
||||
@@ -587,7 +671,6 @@ async fn manages_approval_request_lifecycle() {
|
||||
.claim_next_recoverable_approval_request(
|
||||
timestamp("2026-03-25T12:02:01Z"),
|
||||
timestamp("2026-03-25T12:01:59Z"),
|
||||
timestamp("2026-03-25T11:55:00Z"),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
@@ -608,23 +691,10 @@ async fn manages_approval_request_lifecycle() {
|
||||
.claim_next_recoverable_approval_request(
|
||||
timestamp("2026-03-25T12:02:10Z"),
|
||||
timestamp("2026-03-25T12:02:09Z"),
|
||||
timestamp("2026-03-25T12:01:59Z"),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(fresh_claim.is_none());
|
||||
let recovered = registry
|
||||
.claim_next_recoverable_approval_request(
|
||||
timestamp("2026-03-25T12:03:00Z"),
|
||||
timestamp("2026-03-25T12:02:59Z"),
|
||||
timestamp("2026-03-25T12:02:30Z"),
|
||||
)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(recovered.approval.id, approval.id);
|
||||
assert_eq!(recovered.approval.status, ApprovalRequestStatus::Executing);
|
||||
|
||||
let completed = registry
|
||||
.finish_approval_request(FinishApprovalRequest {
|
||||
workspace_id: &workspace_id,
|
||||
@@ -637,7 +707,6 @@ async fn manages_approval_request_lifecycle() {
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(completed.approval.status, ApprovalRequestStatus::Completed);
|
||||
assert_eq!(
|
||||
completed.approval.response_payload,
|
||||
@@ -653,9 +722,81 @@ async fn manages_approval_request_lifecycle() {
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(completed_by_status.len(), 1);
|
||||
|
||||
let mut interrupted_approval = approval.clone();
|
||||
interrupted_approval.id = ApprovalRequestId::new("approval_interrupted_01");
|
||||
interrupted_approval.request_payload = json!({"amount": 150});
|
||||
interrupted_approval.created_at = timestamp("2026-03-25T12:02:10Z");
|
||||
registry
|
||||
.create_approval_request(CreateApprovalRequest {
|
||||
approval: &interrupted_approval,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
registry
|
||||
.decide_approval_request(DecideApprovalRequest {
|
||||
workspace_id: &workspace_id,
|
||||
agent_id: &agent.id,
|
||||
approval_id: &interrupted_approval.id,
|
||||
status: ApprovalRequestStatus::Approved,
|
||||
decided_at: timestamp("2026-03-25T12:02:20Z"),
|
||||
decided_by_key_id: &approval_key.id,
|
||||
response_payload: Some(json!({"approve": "yes"})),
|
||||
decision_note: None,
|
||||
})
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
registry
|
||||
.claim_approval_request(
|
||||
&workspace_id,
|
||||
&agent.id,
|
||||
&interrupted_approval.id,
|
||||
timestamp("2026-03-25T12:02:21Z"),
|
||||
)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
|
||||
let recovered = registry
|
||||
.claim_next_recoverable_approval_request(
|
||||
timestamp("2026-03-25T12:03:00Z"),
|
||||
timestamp("2026-03-25T12:02:59Z"),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(recovered.is_none());
|
||||
|
||||
let interrupted = registry
|
||||
.fail_next_interrupted_approval_request(timestamp("2026-03-25T12:02:30Z"))
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(interrupted.approval.id, interrupted_approval.id);
|
||||
assert_eq!(interrupted.approval.status, ApprovalRequestStatus::Failed);
|
||||
assert_eq!(
|
||||
completed_by_status[0].approval.status,
|
||||
ApprovalRequestStatus::Completed
|
||||
interrupted.approval.response_payload,
|
||||
Some(json!({
|
||||
"error": {
|
||||
"code": "approval_execution_outcome_unknown",
|
||||
"message": "execution was interrupted; the operation was not retried automatically"
|
||||
}
|
||||
}))
|
||||
);
|
||||
|
||||
let failed_by_status = registry
|
||||
.list_approval_requests(ListApprovalRequestsQuery {
|
||||
workspace_id: &workspace_id,
|
||||
status: Some(ApprovalRequestStatus::Failed),
|
||||
limit: 10,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(failed_by_status.len(), 1);
|
||||
assert_eq!(
|
||||
failed_by_status[0].approval.status,
|
||||
ApprovalRequestStatus::Failed
|
||||
);
|
||||
|
||||
let pending_after_decision = registry
|
||||
|
||||
@@ -3,6 +3,7 @@ name = "crank-runtime"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
rust-version.workspace = true
|
||||
publish.workspace = true
|
||||
version.workspace = true
|
||||
|
||||
[lib]
|
||||
@@ -19,19 +20,22 @@ crank-adapter-rest = { path = "../crank-adapter-rest" }
|
||||
crank-core = { path = "../crank-core" }
|
||||
crank-mapping = { path = "../crank-mapping" }
|
||||
crank-schema = { path = "../crank-schema" }
|
||||
crank-trace = { path = "../crank-trace" }
|
||||
hkdf.workspace = true
|
||||
metrics.workspace = true
|
||||
redis = { version = "0.29", features = ["tokio-comp", "connection-manager"] }
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
sha2.workspace = true
|
||||
thiserror.workspace = true
|
||||
time.workspace = true
|
||||
tokio = { workspace = true, features = ["sync"] }
|
||||
tokio = { workspace = true, features = ["sync", "time"] }
|
||||
tracing.workspace = true
|
||||
uuid.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
axum.workspace = true
|
||||
futures-util = "0.3"
|
||||
testcontainers.workspace = true
|
||||
time.workspace = true
|
||||
tracing-subscriber.workspace = true
|
||||
|
||||
@@ -8,9 +8,9 @@ use std::{
|
||||
|
||||
use async_trait::async_trait;
|
||||
use crank_core::{
|
||||
CacheBackend, CacheScope, CacheStoreError, CachedResponse, CoordinationStateStore,
|
||||
CoordinationStateValue, RateLimitBucketState, RateLimitStateStore, ReplayGuardStatus,
|
||||
ReplayGuardStore, ResponseCacheStore,
|
||||
CacheBackend, CacheScope, CacheStoreError, CachedResponse, CoordinationStateReservation,
|
||||
CoordinationStateStore, CoordinationStateValue, RateLimitBucketState, RateLimitDecision,
|
||||
RateLimitStateStore, ReplayGuardStatus, ReplayGuardStore, ResponseCacheStore,
|
||||
};
|
||||
use redis::{Client, aio::ConnectionManager};
|
||||
use serde::de::DeserializeOwned;
|
||||
@@ -307,6 +307,43 @@ impl RateLimitStateStore for InMemoryRateLimitStateStore {
|
||||
entries.remove(key);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn consume_token(
|
||||
&self,
|
||||
key: &str,
|
||||
burst_tokens_micros: u64,
|
||||
refill_per_second_micros: u64,
|
||||
now_unix_ms: i64,
|
||||
ttl: Duration,
|
||||
) -> Result<RateLimitDecision, CacheStoreError> {
|
||||
validate_key(key)?;
|
||||
validate_rate_limit_parameters(burst_tokens_micros, refill_per_second_micros)?;
|
||||
let expires_at = expiry_from_ttl(ttl)?;
|
||||
let now = Instant::now();
|
||||
let mut entries = self.entries.write().await;
|
||||
retain_unexpired(&mut entries, now);
|
||||
let state = entries
|
||||
.get(key)
|
||||
.map(|entry| entry.value)
|
||||
.unwrap_or(RateLimitBucketState {
|
||||
tokens_micros: burst_tokens_micros,
|
||||
last_refill_unix_ms: now_unix_ms,
|
||||
});
|
||||
let (state, decision) = consume_bucket_token(
|
||||
state,
|
||||
burst_tokens_micros,
|
||||
refill_per_second_micros,
|
||||
now_unix_ms,
|
||||
);
|
||||
entries.insert(
|
||||
key.to_owned(),
|
||||
ExpiringValue {
|
||||
value: state,
|
||||
expires_at,
|
||||
},
|
||||
);
|
||||
Ok(decision)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -373,6 +410,64 @@ impl CoordinationStateStore for InMemoryCoordinationStateStore {
|
||||
entries.remove(&storage_key);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn take_value(
|
||||
&self,
|
||||
scope: CacheScope,
|
||||
key: &str,
|
||||
) -> Result<Option<CoordinationStateValue>, CacheStoreError> {
|
||||
validate_key(key)?;
|
||||
let storage_key = scoped_key(scope, key);
|
||||
let now = Instant::now();
|
||||
let mut entries = self.entries.write().await;
|
||||
retain_unexpired(&mut entries, now);
|
||||
Ok(entries.remove(&storage_key).map(|entry| entry.value))
|
||||
}
|
||||
|
||||
async fn reserve_value(
|
||||
&self,
|
||||
scope: CacheScope,
|
||||
key: &str,
|
||||
value: CoordinationStateValue,
|
||||
ttl: Duration,
|
||||
) -> Result<CoordinationStateReservation, CacheStoreError> {
|
||||
validate_key(key)?;
|
||||
let expires_at = expiry_from_ttl(ttl)?;
|
||||
let storage_key = scoped_key(scope, key);
|
||||
let now = Instant::now();
|
||||
let mut entries = self.entries.write().await;
|
||||
retain_unexpired(&mut entries, now);
|
||||
if let Some(existing) = entries.get(&storage_key) {
|
||||
return Ok(CoordinationStateReservation::Existing(
|
||||
existing.value.clone(),
|
||||
));
|
||||
}
|
||||
entries.insert(storage_key, ExpiringValue { value, expires_at });
|
||||
Ok(CoordinationStateReservation::Reserved)
|
||||
}
|
||||
|
||||
async fn compare_and_set_value(
|
||||
&self,
|
||||
scope: CacheScope,
|
||||
key: &str,
|
||||
expected: &CoordinationStateValue,
|
||||
value: CoordinationStateValue,
|
||||
ttl: Duration,
|
||||
) -> Result<bool, CacheStoreError> {
|
||||
validate_key(key)?;
|
||||
let expires_at = expiry_from_ttl(ttl)?;
|
||||
let storage_key = scoped_key(scope, key);
|
||||
let now = Instant::now();
|
||||
let mut entries = self.entries.write().await;
|
||||
retain_unexpired(&mut entries, now);
|
||||
let matches = entries
|
||||
.get(&storage_key)
|
||||
.is_some_and(|entry| entry.value == *expected);
|
||||
if matches {
|
||||
entries.insert(storage_key, ExpiringValue { value, expires_at });
|
||||
}
|
||||
Ok(matches)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -413,6 +508,65 @@ impl RateLimitStateStore for RedisCacheStore {
|
||||
async fn delete_bucket(&self, key: &str) -> Result<(), CacheStoreError> {
|
||||
self.delete_kind_key("rate_limit", key).await
|
||||
}
|
||||
|
||||
async fn consume_token(
|
||||
&self,
|
||||
key: &str,
|
||||
burst_tokens_micros: u64,
|
||||
refill_per_second_micros: u64,
|
||||
now_unix_ms: i64,
|
||||
ttl: Duration,
|
||||
) -> Result<RateLimitDecision, CacheStoreError> {
|
||||
validate_key(key)?;
|
||||
validate_rate_limit_parameters(burst_tokens_micros, refill_per_second_micros)?;
|
||||
let storage_key = self.prefixed_key("rate_limit", key);
|
||||
let ttl_ms = self.ttl_ms(ttl)?;
|
||||
let mut connection = self.connection_manager.clone();
|
||||
let script = r#"
|
||||
local current = redis.call('GET', KEYS[1])
|
||||
local tokens = tonumber(ARGV[1])
|
||||
local last_refill = tonumber(ARGV[3])
|
||||
if current then
|
||||
local decoded = cjson.decode(current)
|
||||
tokens = tonumber(decoded.tokens_micros)
|
||||
last_refill = tonumber(decoded.last_refill_unix_ms)
|
||||
end
|
||||
local effective_now = math.max(tonumber(ARGV[3]), last_refill)
|
||||
local elapsed = effective_now - last_refill
|
||||
local replenished = math.floor(elapsed * tonumber(ARGV[2]) / 1000)
|
||||
tokens = math.min(tonumber(ARGV[1]), tokens + replenished)
|
||||
local allowed = 0
|
||||
local retry_after = 0
|
||||
if tokens >= 1000000 then
|
||||
tokens = tokens - 1000000
|
||||
allowed = 1
|
||||
else
|
||||
local missing = 1000000 - tokens
|
||||
retry_after = math.max(1, math.ceil(missing * 1000 / tonumber(ARGV[2])))
|
||||
end
|
||||
redis.call('PSETEX', KEYS[1], ARGV[4], cjson.encode({
|
||||
tokens_micros = tokens,
|
||||
last_refill_unix_ms = effective_now
|
||||
}))
|
||||
return {allowed, retry_after}
|
||||
"#;
|
||||
let (allowed, retry_after_ms): (u8, u64) = redis::cmd("EVAL")
|
||||
.arg(script)
|
||||
.arg(1)
|
||||
.arg(storage_key)
|
||||
.arg(burst_tokens_micros)
|
||||
.arg(refill_per_second_micros)
|
||||
.arg(now_unix_ms)
|
||||
.arg(ttl_ms)
|
||||
.query_async(&mut connection)
|
||||
.await
|
||||
.map_err(|source| self.unavailable(source))?;
|
||||
Ok(if allowed == 1 {
|
||||
RateLimitDecision::Allowed
|
||||
} else {
|
||||
RateLimitDecision::Rejected { retry_after_ms }
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -472,6 +626,126 @@ impl CoordinationStateStore for RedisCacheStore {
|
||||
self.delete_kind_key(Self::coordination_kind(scope), key)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn take_value(
|
||||
&self,
|
||||
scope: CacheScope,
|
||||
key: &str,
|
||||
) -> Result<Option<CoordinationStateValue>, CacheStoreError> {
|
||||
validate_key(key)?;
|
||||
let storage_key = self.prefixed_key(Self::coordination_kind(scope), key);
|
||||
let mut connection = self.connection_manager.clone();
|
||||
let encoded: Option<Vec<u8>> = redis::cmd("EVAL")
|
||||
.arg("local value = redis.call('GET', KEYS[1]); if value then redis.call('DEL', KEYS[1]); end; return value")
|
||||
.arg(1)
|
||||
.arg(storage_key)
|
||||
.query_async(&mut connection)
|
||||
.await
|
||||
.map_err(|source| self.unavailable(source))?;
|
||||
encoded
|
||||
.map(|bytes| self.deserialize_value(&bytes))
|
||||
.transpose()
|
||||
}
|
||||
|
||||
async fn reserve_value(
|
||||
&self,
|
||||
scope: CacheScope,
|
||||
key: &str,
|
||||
value: CoordinationStateValue,
|
||||
ttl: Duration,
|
||||
) -> Result<CoordinationStateReservation, CacheStoreError> {
|
||||
validate_key(key)?;
|
||||
let storage_key = self.prefixed_key(Self::coordination_kind(scope), key);
|
||||
let encoded = self.serialize_value(&value)?;
|
||||
let ttl_ms = self.ttl_ms(ttl)?;
|
||||
let mut connection = self.connection_manager.clone();
|
||||
let (reserved, existing): (u8, Vec<u8>) = redis::cmd("EVAL")
|
||||
.arg("local current = redis.call('GET', KEYS[1]); if current then return {0, current}; end; redis.call('PSETEX', KEYS[1], ARGV[2], ARGV[1]); return {1, ''}")
|
||||
.arg(1)
|
||||
.arg(storage_key)
|
||||
.arg(encoded)
|
||||
.arg(ttl_ms)
|
||||
.query_async(&mut connection)
|
||||
.await
|
||||
.map_err(|source| self.unavailable(source))?;
|
||||
if reserved == 1 {
|
||||
Ok(CoordinationStateReservation::Reserved)
|
||||
} else {
|
||||
Ok(CoordinationStateReservation::Existing(
|
||||
self.deserialize_value(&existing)?,
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
async fn compare_and_set_value(
|
||||
&self,
|
||||
scope: CacheScope,
|
||||
key: &str,
|
||||
expected: &CoordinationStateValue,
|
||||
value: CoordinationStateValue,
|
||||
ttl: Duration,
|
||||
) -> Result<bool, CacheStoreError> {
|
||||
validate_key(key)?;
|
||||
let storage_key = self.prefixed_key(Self::coordination_kind(scope), key);
|
||||
let expected = self.serialize_value(expected)?;
|
||||
let value = self.serialize_value(&value)?;
|
||||
let ttl_ms = self.ttl_ms(ttl)?;
|
||||
let mut connection = self.connection_manager.clone();
|
||||
let replaced: u8 = redis::cmd("EVAL")
|
||||
.arg("local current = redis.call('GET', KEYS[1]); if current ~= ARGV[1] then return 0; end; redis.call('PSETEX', KEYS[1], ARGV[3], ARGV[2]); return 1")
|
||||
.arg(1)
|
||||
.arg(storage_key)
|
||||
.arg(expected)
|
||||
.arg(value)
|
||||
.arg(ttl_ms)
|
||||
.query_async(&mut connection)
|
||||
.await
|
||||
.map_err(|source| self.unavailable(source))?;
|
||||
Ok(replaced == 1)
|
||||
}
|
||||
}
|
||||
|
||||
fn consume_bucket_token(
|
||||
mut state: RateLimitBucketState,
|
||||
burst_tokens_micros: u64,
|
||||
refill_per_second_micros: u64,
|
||||
now_unix_ms: i64,
|
||||
) -> (RateLimitBucketState, RateLimitDecision) {
|
||||
let effective_now = now_unix_ms.max(state.last_refill_unix_ms);
|
||||
let elapsed_ms =
|
||||
u64::try_from(effective_now.saturating_sub(state.last_refill_unix_ms)).unwrap_or(u64::MAX);
|
||||
let replenished =
|
||||
u128::from(elapsed_ms).saturating_mul(u128::from(refill_per_second_micros)) / 1000;
|
||||
state.tokens_micros = u128::from(state.tokens_micros)
|
||||
.saturating_add(replenished)
|
||||
.min(u128::from(burst_tokens_micros))
|
||||
.try_into()
|
||||
.unwrap_or(burst_tokens_micros);
|
||||
state.last_refill_unix_ms = effective_now;
|
||||
if state.tokens_micros >= 1_000_000 {
|
||||
state.tokens_micros -= 1_000_000;
|
||||
return (state, RateLimitDecision::Allowed);
|
||||
}
|
||||
let missing = 1_000_000_u64.saturating_sub(state.tokens_micros);
|
||||
let retry_after_ms = u128::from(missing)
|
||||
.saturating_mul(1000)
|
||||
.div_ceil(u128::from(refill_per_second_micros))
|
||||
.max(1)
|
||||
.try_into()
|
||||
.unwrap_or(u64::MAX);
|
||||
(state, RateLimitDecision::Rejected { retry_after_ms })
|
||||
}
|
||||
|
||||
fn validate_rate_limit_parameters(
|
||||
burst_tokens_micros: u64,
|
||||
refill_per_second_micros: u64,
|
||||
) -> Result<(), CacheStoreError> {
|
||||
if burst_tokens_micros < 1_000_000 || refill_per_second_micros == 0 {
|
||||
return Err(CacheStoreError::InvalidKey {
|
||||
message: "rate limit capacity and refill rate must be positive".to_owned(),
|
||||
});
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_key(key: &str) -> Result<(), CacheStoreError> {
|
||||
@@ -530,7 +804,11 @@ fn parse_optional_string(name: &'static str) -> Result<Option<String>, RuntimeCa
|
||||
fn parse_optional_u64(name: &'static str) -> Result<Option<u64>, RuntimeCacheConfigError> {
|
||||
match env::var(name) {
|
||||
Ok(raw) => {
|
||||
let value = raw
|
||||
let trimmed = raw.trim();
|
||||
if trimmed.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
let value = trimmed
|
||||
.parse::<u64>()
|
||||
.map_err(|source| RuntimeCacheConfigError::InvalidTtl { value: raw, source })?;
|
||||
if value == 0 {
|
||||
|
||||
@@ -50,6 +50,12 @@ pub async fn confirm_operation(
|
||||
consume_confirmation_token(store, &scope, provided_token, &input_hash).await
|
||||
}
|
||||
|
||||
pub(crate) fn is_applicable(operation: &RuntimeOperation) -> bool {
|
||||
effective_safety_policy(operation)
|
||||
.class
|
||||
.requires_confirmation()
|
||||
}
|
||||
|
||||
fn effective_safety_policy(operation: &RuntimeOperation) -> OperationSafetyPolicy {
|
||||
operation
|
||||
.execution_config
|
||||
@@ -137,13 +143,12 @@ async fn consume_confirmation_token(
|
||||
) -> Result<(), RuntimeError> {
|
||||
let key = confirmation_cache_key(operation_scope, token);
|
||||
let stored = store
|
||||
.get_value(CacheScope::Coordination, &key)
|
||||
.take_value(CacheScope::Coordination, &key)
|
||||
.await
|
||||
.map_err(|error| RuntimeError::InvalidPreparedRequest {
|
||||
field: "confirmation_token".to_owned(),
|
||||
reason: error.to_string(),
|
||||
})?;
|
||||
let _ = store.delete_value(CacheScope::Coordination, &key).await;
|
||||
|
||||
let Some(stored) = stored else {
|
||||
return Err(RuntimeError::InvalidConfirmationToken {
|
||||
|
||||
@@ -38,6 +38,14 @@ pub enum RuntimeError {
|
||||
InvalidConfirmationToken { operation_id: String },
|
||||
#[error("confirmation store is unavailable for operation {operation_id}")]
|
||||
ConfirmationStoreUnavailable { operation_id: String },
|
||||
#[error("idempotency store is unavailable for operation {operation_id}")]
|
||||
IdempotencyStoreUnavailable { operation_id: String },
|
||||
#[error("operation {operation_id} is already executing for this idempotency key")]
|
||||
IdempotencyInProgress { operation_id: String },
|
||||
#[error("idempotency key for operation {operation_id} was reused with different input")]
|
||||
IdempotencyConflict { operation_id: String },
|
||||
#[error("the outcome of operation {operation_id} is unknown; automatic retry is unsafe")]
|
||||
IdempotencyOutcomeUnknown { operation_id: String },
|
||||
#[error("auth profile {auth_profile_id} was not found")]
|
||||
MissingAuthProfile { auth_profile_id: String },
|
||||
#[error("secret {secret_id} was not found")]
|
||||
|
||||
@@ -1,14 +1,16 @@
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
use std::time::Instant;
|
||||
|
||||
use crank_core::{
|
||||
AdapterRegistry, CoordinationStateStore, ExecutionMode, InvocationStatus, MeteringEvent,
|
||||
ResponseCacheStore, SharedMeteringSink, SharedProtocolAdapter,
|
||||
AdapterRegistry, CoordinationStateStore, ExecutionMode, InvocationSource, InvocationStatus,
|
||||
MeteringEvent, ResponseCacheStore, SharedMeteringSink, SharedProtocolAdapter,
|
||||
};
|
||||
use crank_trace::{ErrorCategory, Stage, StageOutcome};
|
||||
use metrics::Gauge;
|
||||
use serde_json::{Map, Value, json};
|
||||
use time::OffsetDateTime;
|
||||
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
|
||||
use tracing::debug;
|
||||
use tracing::{Instrument, Span, debug};
|
||||
|
||||
use crate::{
|
||||
AdapterResponse, PreparedRequest, ResolvedAuth, RuntimeError, RuntimeLimits, RuntimeOperation,
|
||||
@@ -176,18 +178,31 @@ impl RuntimeExecutor {
|
||||
request: RuntimeExecutionRequest<'_>,
|
||||
) -> Result<Value, RuntimeError> {
|
||||
log_runtime_event("unary.execute", request.operation, request.request_context);
|
||||
let _permit = self.acquire_unary_permit(request.operation)?;
|
||||
let started_at = Instant::now();
|
||||
let prepared_request = self.prepare_request(request.operation, request.input)?;
|
||||
let prepared_request = apply_resolved_auth(prepared_request, request.resolved_auth);
|
||||
let result = self
|
||||
.execute_prepared(
|
||||
let runtime_span = Stage::RuntimeExecute.span();
|
||||
let result = async {
|
||||
let _permit = self.acquire_unary_permit(request.operation)?;
|
||||
let _inflight = RuntimeInFlightGuard::new();
|
||||
let mapping_span = Stage::RuntimeArgumentsMap.span();
|
||||
let prepared_request =
|
||||
mapping_span.in_scope(|| self.prepare_request(request.operation, request.input));
|
||||
record_runtime_result(&mapping_span, &prepared_request);
|
||||
drop(mapping_span);
|
||||
let prepared_request = prepared_request?;
|
||||
let prepared_request = apply_resolved_auth(prepared_request, request.resolved_auth);
|
||||
self.execute_prepared(
|
||||
request.operation,
|
||||
request.input,
|
||||
prepared_request,
|
||||
request.request_context,
|
||||
)
|
||||
.await;
|
||||
.await
|
||||
}
|
||||
.instrument(runtime_span.clone())
|
||||
.await;
|
||||
record_runtime_result(&runtime_span, &result);
|
||||
drop(runtime_span);
|
||||
record_execution_metrics(request.request_context, &result, started_at);
|
||||
self.record_metering(
|
||||
request.operation,
|
||||
request.request_context,
|
||||
@@ -221,59 +236,132 @@ impl RuntimeExecutor {
|
||||
request_context: Option<&RuntimeRequestContext>,
|
||||
) -> Result<Value, RuntimeError> {
|
||||
let mut prepared_request = prepared_request;
|
||||
let idempotency_key =
|
||||
crate::idempotency::prepare_idempotency(operation, input, &mut prepared_request)?;
|
||||
crate::confirmation::confirm_operation(
|
||||
self.coordination_store.as_deref(),
|
||||
let idempotency_applicable = crate::idempotency::policy(operation).is_some();
|
||||
let idempotency_key = match crate::idempotency::prepare_idempotency(
|
||||
operation,
|
||||
input,
|
||||
request_context,
|
||||
)
|
||||
.await?;
|
||||
if let Some(response) = self
|
||||
.load_idempotent_adapter_response(
|
||||
&mut prepared_request,
|
||||
) {
|
||||
Ok(key) => key,
|
||||
Err(error) if idempotency_applicable => {
|
||||
let span = Stage::RuntimeIdempotency.span();
|
||||
StageOutcome::Error.record(&span);
|
||||
ErrorCategory::Idempotency.record(&span);
|
||||
return Err(error);
|
||||
}
|
||||
Err(error) => return Err(error),
|
||||
};
|
||||
|
||||
if crate::confirmation::is_applicable(operation) {
|
||||
let approval_span = Stage::ApprovalCheck.span();
|
||||
let approval_result = crate::confirmation::confirm_operation(
|
||||
self.coordination_store.as_deref(),
|
||||
operation,
|
||||
input,
|
||||
request_context,
|
||||
)
|
||||
.instrument(approval_span.clone())
|
||||
.await;
|
||||
match &approval_result {
|
||||
Ok(()) => StageOutcome::Success.record(&approval_span),
|
||||
Err(RuntimeError::ConfirmationRequired { .. }) => {
|
||||
StageOutcome::Required.record(&approval_span);
|
||||
ErrorCategory::Approval.record(&approval_span);
|
||||
}
|
||||
Err(error) => {
|
||||
StageOutcome::Error.record(&approval_span);
|
||||
runtime_error_category(error).record(&approval_span);
|
||||
}
|
||||
}
|
||||
drop(approval_span);
|
||||
approval_result?;
|
||||
}
|
||||
|
||||
let idempotency = if idempotency_applicable {
|
||||
let idempotency_span = Stage::RuntimeIdempotency.span();
|
||||
let result = crate::idempotency::begin(
|
||||
self.coordination_store.as_deref(),
|
||||
operation,
|
||||
input,
|
||||
idempotency_key.as_deref(),
|
||||
request_context,
|
||||
)
|
||||
.await
|
||||
{
|
||||
let finalized_output = finalize_output(operation, &response)?;
|
||||
operation.output_schema.validate_shape(&finalized_output)?;
|
||||
return Ok(finalized_output);
|
||||
.instrument(idempotency_span.clone())
|
||||
.await;
|
||||
match &result {
|
||||
Ok(crate::idempotency::IdempotencyAction::Execute(_)) => {
|
||||
StageOutcome::Execute.record(&idempotency_span);
|
||||
}
|
||||
Ok(crate::idempotency::IdempotencyAction::Replay(_)) => {
|
||||
StageOutcome::Replay.record(&idempotency_span);
|
||||
}
|
||||
Ok(crate::idempotency::IdempotencyAction::Disabled) => {
|
||||
StageOutcome::Skipped.record(&idempotency_span);
|
||||
}
|
||||
Err(error) => {
|
||||
StageOutcome::Error.record(&idempotency_span);
|
||||
runtime_error_category(error).record(&idempotency_span);
|
||||
}
|
||||
}
|
||||
drop(idempotency_span);
|
||||
result?
|
||||
} else {
|
||||
crate::idempotency::IdempotencyAction::Disabled
|
||||
};
|
||||
if let crate::idempotency::IdempotencyAction::Replay(response) = &idempotency {
|
||||
return transform_response(operation, response);
|
||||
}
|
||||
|
||||
let adapter_response = match self
|
||||
let adapter_result = match self
|
||||
.load_cached_adapter_response(operation, &prepared_request, request_context)
|
||||
.await
|
||||
{
|
||||
Some(response) => response,
|
||||
Some(response) => Ok(response),
|
||||
None => {
|
||||
let adapter_response = self
|
||||
.execute_adapter(operation, prepared_request.clone(), request_context)
|
||||
.await?;
|
||||
self.store_cached_adapter_response(
|
||||
operation,
|
||||
&prepared_request,
|
||||
request_context,
|
||||
&adapter_response,
|
||||
)
|
||||
.await;
|
||||
self.store_idempotent_adapter_response(
|
||||
operation,
|
||||
idempotency_key.as_deref(),
|
||||
request_context,
|
||||
&adapter_response,
|
||||
)
|
||||
.await;
|
||||
.await;
|
||||
if let Ok(response) = &adapter_response {
|
||||
self.store_cached_adapter_response(
|
||||
operation,
|
||||
&prepared_request,
|
||||
request_context,
|
||||
response,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
adapter_response
|
||||
}
|
||||
};
|
||||
let finalized_output = finalize_output(operation, &adapter_response)?;
|
||||
|
||||
operation.output_schema.validate_shape(&finalized_output)?;
|
||||
|
||||
Ok(finalized_output)
|
||||
let adapter_response = match adapter_result {
|
||||
Ok(response) => response,
|
||||
Err(error) => {
|
||||
if let crate::idempotency::IdempotencyAction::Execute(reservation) = &idempotency
|
||||
&& let Some(store) = self.coordination_store.as_deref()
|
||||
{
|
||||
let idempotency_span = Stage::RuntimeIdempotency.span();
|
||||
let cleanup_result =
|
||||
crate::idempotency::mark_outcome_unknown(store, operation, reservation)
|
||||
.instrument(idempotency_span.clone())
|
||||
.await;
|
||||
record_runtime_result(&idempotency_span, &cleanup_result);
|
||||
}
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
if let crate::idempotency::IdempotencyAction::Execute(reservation) = &idempotency
|
||||
&& let Some(store) = self.coordination_store.as_deref()
|
||||
{
|
||||
let idempotency_span = Stage::RuntimeIdempotency.span();
|
||||
let completion_result =
|
||||
crate::idempotency::complete(store, operation, reservation, &adapter_response)
|
||||
.instrument(idempotency_span.clone())
|
||||
.await;
|
||||
record_runtime_result(&idempotency_span, &completion_result);
|
||||
drop(idempotency_span);
|
||||
completion_result?;
|
||||
}
|
||||
transform_response(operation, &adapter_response)
|
||||
}
|
||||
|
||||
async fn record_metering<T>(
|
||||
@@ -323,7 +411,6 @@ impl RuntimeExecutor {
|
||||
let prepared_request = adapter_prepared_request(
|
||||
operation,
|
||||
&prepared_request,
|
||||
request_context,
|
||||
operation.execution_config.timeout_ms,
|
||||
);
|
||||
let adapter_context = adapter_request_context(request_context);
|
||||
@@ -360,11 +447,11 @@ impl RuntimeExecutor {
|
||||
let cache_key = response_cache_key(operation, prepared_request, request_context)?;
|
||||
let cached = match response_cache.get(&cache_key).await {
|
||||
Ok(cached) => cached?,
|
||||
Err(error) => {
|
||||
Err(_) => {
|
||||
debug!(
|
||||
operation_id = %operation.operation_id,
|
||||
cache_key,
|
||||
error = %error,
|
||||
name: "runtime.response_cache.read_failed",
|
||||
operation_id = operation.operation_id.as_str(),
|
||||
error_category = "response_cache",
|
||||
"response cache lookup skipped"
|
||||
);
|
||||
return None;
|
||||
@@ -373,11 +460,11 @@ impl RuntimeExecutor {
|
||||
|
||||
match adapter_response_from_cached(cached) {
|
||||
Ok(response) => Some(response),
|
||||
Err(error) => {
|
||||
Err(_) => {
|
||||
debug!(
|
||||
operation_id = %operation.operation_id,
|
||||
cache_key,
|
||||
error,
|
||||
name: "runtime.response_cache.decode_failed",
|
||||
operation_id = operation.operation_id.as_str(),
|
||||
error_category = "cached_response",
|
||||
"cached response payload was invalid"
|
||||
);
|
||||
let _ = response_cache.delete(&cache_key).await;
|
||||
@@ -412,87 +499,19 @@ impl RuntimeExecutor {
|
||||
return;
|
||||
};
|
||||
|
||||
if let Err(error) = response_cache
|
||||
if response_cache
|
||||
.put(&cache_key, cached_response, cache_ttl)
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
debug!(
|
||||
operation_id = %operation.operation_id,
|
||||
cache_key,
|
||||
error = %error,
|
||||
name: "runtime.response_cache.write_failed",
|
||||
operation_id = operation.operation_id.as_str(),
|
||||
error_category = "response_cache",
|
||||
"response cache write skipped"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
async fn load_idempotent_adapter_response(
|
||||
&self,
|
||||
operation: &RuntimeOperation,
|
||||
idempotency_key: Option<&str>,
|
||||
request_context: Option<&RuntimeRequestContext>,
|
||||
) -> Option<AdapterResponse> {
|
||||
let response_cache = self.response_cache.as_ref()?;
|
||||
let cache_key =
|
||||
crate::idempotency::cache_key(operation, idempotency_key?, request_context)?;
|
||||
let cached = match response_cache.get(&cache_key).await {
|
||||
Ok(cached) => cached?,
|
||||
Err(error) => {
|
||||
debug!(
|
||||
operation_id = %operation.operation_id,
|
||||
cache_key,
|
||||
error = %error,
|
||||
"idempotency cache lookup skipped"
|
||||
);
|
||||
return None;
|
||||
}
|
||||
};
|
||||
|
||||
adapter_response_from_cached(cached).ok()
|
||||
}
|
||||
|
||||
async fn store_idempotent_adapter_response(
|
||||
&self,
|
||||
operation: &RuntimeOperation,
|
||||
idempotency_key: Option<&str>,
|
||||
request_context: Option<&RuntimeRequestContext>,
|
||||
adapter_response: &AdapterResponse,
|
||||
) {
|
||||
let Some(response_cache) = self.response_cache.as_ref() else {
|
||||
return;
|
||||
};
|
||||
if !(200..=299).contains(&adapter_response.status_code) {
|
||||
return;
|
||||
}
|
||||
let Some(policy) = crate::idempotency::policy(operation) else {
|
||||
return;
|
||||
};
|
||||
let Some(cache_key) = crate::idempotency::cache_key(
|
||||
operation,
|
||||
idempotency_key.unwrap_or_default(),
|
||||
request_context,
|
||||
) else {
|
||||
return;
|
||||
};
|
||||
let Some(cached_response) = cached_response_from_adapter(adapter_response) else {
|
||||
return;
|
||||
};
|
||||
|
||||
if let Err(error) = response_cache
|
||||
.put(
|
||||
&cache_key,
|
||||
cached_response,
|
||||
Duration::from_millis(policy.ttl_ms),
|
||||
)
|
||||
.await
|
||||
{
|
||||
debug!(
|
||||
operation_id = %operation.operation_id,
|
||||
cache_key,
|
||||
error = %error,
|
||||
"idempotency cache write skipped"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn adapter_request_context(
|
||||
@@ -552,14 +571,149 @@ fn finalize_output(
|
||||
.unwrap_or_else(|| Value::Object(Map::new())))
|
||||
}
|
||||
|
||||
fn transform_response(
|
||||
operation: &RuntimeOperation,
|
||||
response: &AdapterResponse,
|
||||
) -> Result<Value, RuntimeError> {
|
||||
let span = Stage::RuntimeResponseTransform.span();
|
||||
let result = span.in_scope(|| {
|
||||
let finalized_output = finalize_output(operation, response)?;
|
||||
operation.output_schema.validate_shape(&finalized_output)?;
|
||||
Ok(finalized_output)
|
||||
});
|
||||
match &result {
|
||||
Ok(_) => StageOutcome::Success.record(&span),
|
||||
Err(_) => {
|
||||
StageOutcome::Error.record(&span);
|
||||
ErrorCategory::Transformation.record(&span);
|
||||
}
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
fn record_runtime_result<T>(span: &Span, result: &Result<T, RuntimeError>) {
|
||||
match result {
|
||||
Ok(_) => StageOutcome::Success.record(span),
|
||||
Err(error) => {
|
||||
StageOutcome::Error.record(span);
|
||||
runtime_error_category(error).record(span);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn runtime_error_category(error: &RuntimeError) -> ErrorCategory {
|
||||
match error {
|
||||
RuntimeError::Schema(_) => ErrorCategory::Schema,
|
||||
RuntimeError::Mapping(_) | RuntimeError::InvalidPreparedRequest { .. } => {
|
||||
ErrorCategory::Mapping
|
||||
}
|
||||
RuntimeError::RestAdapter(_)
|
||||
| RuntimeError::ProtocolAdapter(_)
|
||||
| RuntimeError::UnsupportedProtocol { .. }
|
||||
| RuntimeError::UnsupportedExecutionMode { .. } => ErrorCategory::Upstream,
|
||||
RuntimeError::ConcurrencyLimitExceeded { .. } => ErrorCategory::Concurrency,
|
||||
RuntimeError::ConfirmationRequired { .. }
|
||||
| RuntimeError::InvalidConfirmationToken { .. }
|
||||
| RuntimeError::ConfirmationStoreUnavailable { .. } => ErrorCategory::Approval,
|
||||
RuntimeError::IdempotencyStoreUnavailable { .. }
|
||||
| RuntimeError::IdempotencyInProgress { .. }
|
||||
| RuntimeError::IdempotencyConflict { .. }
|
||||
| RuntimeError::IdempotencyOutcomeUnknown { .. } => ErrorCategory::Idempotency,
|
||||
RuntimeError::MissingAuthProfile { .. }
|
||||
| RuntimeError::MissingSecret { .. }
|
||||
| RuntimeError::MissingSecretVersion { .. }
|
||||
| RuntimeError::InvalidAuthSecretValue { .. }
|
||||
| RuntimeError::SecretCrypto { .. } => ErrorCategory::Configuration,
|
||||
}
|
||||
}
|
||||
|
||||
fn try_acquire_limit(
|
||||
limiter: Arc<Semaphore>,
|
||||
kind: &'static str,
|
||||
limit: usize,
|
||||
) -> Result<OwnedSemaphorePermit, RuntimeError> {
|
||||
limiter
|
||||
.try_acquire_owned()
|
||||
.map_err(|_| RuntimeError::ConcurrencyLimitExceeded { kind, limit })
|
||||
limiter.try_acquire_owned().map_err(|_| {
|
||||
metrics::counter!(
|
||||
"crank_runtime_limit_rejections_total",
|
||||
"stage" => "concurrency"
|
||||
)
|
||||
.increment(1);
|
||||
RuntimeError::ConcurrencyLimitExceeded { kind, limit }
|
||||
})
|
||||
}
|
||||
|
||||
fn record_execution_metrics<T>(
|
||||
request_context: Option<&RuntimeRequestContext>,
|
||||
result: &Result<T, RuntimeError>,
|
||||
started_at: Instant,
|
||||
) {
|
||||
let source = request_context
|
||||
.and_then(RuntimeRequestContext::metering_context)
|
||||
.map_or("internal", |context| match context.source {
|
||||
InvocationSource::AdminTestRun => "admin_test_run",
|
||||
InvocationSource::AgentToolCall => "agent_tool_call",
|
||||
});
|
||||
let (outcome, error_kind) = match result {
|
||||
Ok(_) => ("success", "none"),
|
||||
Err(error) => ("error", runtime_error_kind(error)),
|
||||
};
|
||||
|
||||
metrics::counter!(
|
||||
"crank_tool_invocations_total",
|
||||
"source" => source,
|
||||
"outcome" => outcome,
|
||||
"error_kind" => error_kind
|
||||
)
|
||||
.increment(1);
|
||||
metrics::histogram!(
|
||||
"crank_tool_invocation_duration_seconds",
|
||||
"source" => source,
|
||||
"outcome" => outcome
|
||||
)
|
||||
.record(started_at.elapsed().as_secs_f64());
|
||||
}
|
||||
|
||||
fn runtime_error_kind(error: &RuntimeError) -> &'static str {
|
||||
match error {
|
||||
RuntimeError::Schema(_) => "schema",
|
||||
RuntimeError::Mapping(_) => "mapping",
|
||||
RuntimeError::RestAdapter(_) => "rest_adapter",
|
||||
RuntimeError::ProtocolAdapter(_) => "protocol_adapter",
|
||||
RuntimeError::UnsupportedProtocol { .. } => "unsupported_protocol",
|
||||
RuntimeError::UnsupportedExecutionMode { .. } => "unsupported_execution_mode",
|
||||
RuntimeError::ConcurrencyLimitExceeded { .. } => "concurrency_limit",
|
||||
RuntimeError::InvalidPreparedRequest { .. } => "invalid_prepared_request",
|
||||
RuntimeError::ConfirmationRequired { .. } => "confirmation_required",
|
||||
RuntimeError::InvalidConfirmationToken { .. } => "invalid_confirmation_token",
|
||||
RuntimeError::ConfirmationStoreUnavailable { .. } => "confirmation_store",
|
||||
RuntimeError::IdempotencyStoreUnavailable { .. } => "idempotency_store",
|
||||
RuntimeError::IdempotencyInProgress { .. } => "idempotency_in_progress",
|
||||
RuntimeError::IdempotencyConflict { .. } => "idempotency_conflict",
|
||||
RuntimeError::IdempotencyOutcomeUnknown { .. } => "idempotency_outcome_unknown",
|
||||
RuntimeError::MissingAuthProfile { .. } => "missing_auth_profile",
|
||||
RuntimeError::MissingSecret { .. } => "missing_secret",
|
||||
RuntimeError::MissingSecretVersion { .. } => "missing_secret_version",
|
||||
RuntimeError::InvalidAuthSecretValue { .. } => "invalid_auth_secret",
|
||||
RuntimeError::SecretCrypto { .. } => "secret_crypto",
|
||||
}
|
||||
}
|
||||
|
||||
struct RuntimeInFlightGuard {
|
||||
gauge: Gauge,
|
||||
}
|
||||
|
||||
impl RuntimeInFlightGuard {
|
||||
fn new() -> Self {
|
||||
let gauge = metrics::gauge!("crank_runtime_inflight");
|
||||
gauge.increment(1.0);
|
||||
Self { gauge }
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for RuntimeInFlightGuard {
|
||||
fn drop(&mut self) {
|
||||
self.gauge.decrement(1.0);
|
||||
}
|
||||
}
|
||||
|
||||
fn log_runtime_event(
|
||||
@@ -567,19 +721,26 @@ fn log_runtime_event(
|
||||
operation: &RuntimeOperation,
|
||||
request_context: Option<&RuntimeRequestContext>,
|
||||
) {
|
||||
let request_id = request_context
|
||||
.map(|context| context.request_id.as_str())
|
||||
.unwrap_or_default();
|
||||
let correlation_id = request_context
|
||||
.map(|context| context.correlation_id.as_str())
|
||||
.unwrap_or_default();
|
||||
|
||||
debug!(
|
||||
stage,
|
||||
operation_id = %operation.operation_id,
|
||||
protocol = ?operation.protocol,
|
||||
request_id,
|
||||
correlation_id,
|
||||
"runtime execution"
|
||||
);
|
||||
let protocol = match operation.protocol {
|
||||
crank_core::Protocol::Rest => "rest",
|
||||
};
|
||||
if let Some(context) = request_context {
|
||||
debug!(
|
||||
name: "runtime.execution.stage_reached",
|
||||
stage,
|
||||
operation_id = operation.operation_id.as_str(),
|
||||
protocol,
|
||||
request_id = context.request_id.as_str(),
|
||||
correlation_id = context.correlation_id.as_str(),
|
||||
"runtime execution"
|
||||
);
|
||||
} else {
|
||||
debug!(
|
||||
name: "runtime.execution.stage_reached",
|
||||
stage,
|
||||
operation_id = operation.operation_id.as_str(),
|
||||
protocol,
|
||||
"runtime execution"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,9 +1,32 @@
|
||||
use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
|
||||
use crank_core::{HttpMethod, IdempotencyMode, IdempotencyPolicy, Target};
|
||||
use serde_json::Value;
|
||||
use sha2::{Digest, Sha256};
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use crate::{PreparedRequest, RuntimeError, RuntimeOperation, RuntimeRequestContext};
|
||||
use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
|
||||
use crank_core::{
|
||||
CacheScope, CoordinationStateReservation, CoordinationStateStore, CoordinationStateValue,
|
||||
HttpMethod, IdempotencyMode, IdempotencyPolicy, Target,
|
||||
};
|
||||
use serde_json::{Map, Value, json};
|
||||
use sha2::{Digest, Sha256};
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::{
|
||||
AdapterResponse, PreparedRequest, RuntimeError, RuntimeOperation, RuntimeRequestContext,
|
||||
};
|
||||
|
||||
const POLL_INTERVAL: Duration = Duration::from_millis(10);
|
||||
|
||||
pub(crate) enum IdempotencyAction {
|
||||
Disabled,
|
||||
Execute(IdempotencyReservation),
|
||||
Replay(AdapterResponse),
|
||||
}
|
||||
|
||||
pub(crate) struct IdempotencyReservation {
|
||||
key: String,
|
||||
initial: CoordinationStateValue,
|
||||
fingerprint: String,
|
||||
result_ttl: Duration,
|
||||
}
|
||||
|
||||
pub fn prepare_idempotency(
|
||||
operation: &RuntimeOperation,
|
||||
@@ -39,6 +62,185 @@ pub fn prepare_idempotency(
|
||||
Ok(Some(key))
|
||||
}
|
||||
|
||||
pub(crate) async fn begin(
|
||||
store: Option<&dyn CoordinationStateStore>,
|
||||
operation: &RuntimeOperation,
|
||||
input: &Value,
|
||||
idempotency_key: Option<&str>,
|
||||
request_context: Option<&RuntimeRequestContext>,
|
||||
) -> Result<IdempotencyAction, RuntimeError> {
|
||||
let Some(policy) = policy(operation) else {
|
||||
return Ok(IdempotencyAction::Disabled);
|
||||
};
|
||||
let Some(idempotency_key) = idempotency_key else {
|
||||
return Ok(IdempotencyAction::Disabled);
|
||||
};
|
||||
let Some(key) = cache_key(operation, idempotency_key, request_context) else {
|
||||
return Err(RuntimeError::IdempotencyStoreUnavailable {
|
||||
operation_id: operation.operation_id.as_str().to_owned(),
|
||||
});
|
||||
};
|
||||
let Some(store) = store else {
|
||||
return Err(RuntimeError::IdempotencyStoreUnavailable {
|
||||
operation_id: operation.operation_id.as_str().to_owned(),
|
||||
});
|
||||
};
|
||||
|
||||
let fingerprint = request_fingerprint(input)?;
|
||||
let result_ttl = Duration::from_millis(policy.ttl_ms);
|
||||
let reservation_ttl = result_ttl.max(Duration::from_millis(
|
||||
operation.execution_config.timeout_ms.saturating_add(1_000),
|
||||
));
|
||||
let initial = in_progress_value(&fingerprint);
|
||||
let reservation = store
|
||||
.reserve_value(
|
||||
CacheScope::Coordination,
|
||||
&key,
|
||||
initial.clone(),
|
||||
reservation_ttl,
|
||||
)
|
||||
.await
|
||||
.map_err(|_| RuntimeError::IdempotencyStoreUnavailable {
|
||||
operation_id: operation.operation_id.as_str().to_owned(),
|
||||
})?;
|
||||
|
||||
match reservation {
|
||||
CoordinationStateReservation::Reserved => {
|
||||
Ok(IdempotencyAction::Execute(IdempotencyReservation {
|
||||
key,
|
||||
initial,
|
||||
fingerprint,
|
||||
result_ttl,
|
||||
}))
|
||||
}
|
||||
CoordinationStateReservation::Existing(existing) => {
|
||||
resolve_existing(store, operation, &key, &fingerprint, existing).await
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn complete(
|
||||
store: &dyn CoordinationStateStore,
|
||||
operation: &RuntimeOperation,
|
||||
reservation: &IdempotencyReservation,
|
||||
response: &AdapterResponse,
|
||||
) -> Result<(), RuntimeError> {
|
||||
let completed = CoordinationStateValue {
|
||||
payload: json!({
|
||||
"state": "completed",
|
||||
"fingerprint": reservation.fingerprint,
|
||||
"response": response,
|
||||
}),
|
||||
};
|
||||
let replaced = store
|
||||
.compare_and_set_value(
|
||||
CacheScope::Coordination,
|
||||
&reservation.key,
|
||||
&reservation.initial,
|
||||
completed,
|
||||
reservation.result_ttl,
|
||||
)
|
||||
.await
|
||||
.map_err(|_| RuntimeError::IdempotencyStoreUnavailable {
|
||||
operation_id: operation.operation_id.as_str().to_owned(),
|
||||
})?;
|
||||
if replaced {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(RuntimeError::IdempotencyOutcomeUnknown {
|
||||
operation_id: operation.operation_id.as_str().to_owned(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn mark_outcome_unknown(
|
||||
store: &dyn CoordinationStateStore,
|
||||
operation: &RuntimeOperation,
|
||||
reservation: &IdempotencyReservation,
|
||||
) -> Result<(), RuntimeError> {
|
||||
let unknown = CoordinationStateValue {
|
||||
payload: json!({
|
||||
"state": "outcome_unknown",
|
||||
"fingerprint": reservation.fingerprint,
|
||||
}),
|
||||
};
|
||||
let replaced = store
|
||||
.compare_and_set_value(
|
||||
CacheScope::Coordination,
|
||||
&reservation.key,
|
||||
&reservation.initial,
|
||||
unknown,
|
||||
reservation.result_ttl,
|
||||
)
|
||||
.await
|
||||
.map_err(|_| RuntimeError::IdempotencyStoreUnavailable {
|
||||
operation_id: operation.operation_id.as_str().to_owned(),
|
||||
})?;
|
||||
if replaced {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(RuntimeError::IdempotencyOutcomeUnknown {
|
||||
operation_id: operation.operation_id.as_str().to_owned(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
async fn resolve_existing(
|
||||
store: &dyn CoordinationStateStore,
|
||||
operation: &RuntimeOperation,
|
||||
key: &str,
|
||||
fingerprint: &str,
|
||||
mut existing: CoordinationStateValue,
|
||||
) -> Result<IdempotencyAction, RuntimeError> {
|
||||
let operation_id = operation.operation_id.as_str().to_owned();
|
||||
let deadline =
|
||||
Instant::now() + Duration::from_millis(operation.execution_config.timeout_ms.max(1));
|
||||
loop {
|
||||
let existing_fingerprint = existing.payload.get("fingerprint").and_then(Value::as_str);
|
||||
if existing_fingerprint != Some(fingerprint) {
|
||||
return Err(RuntimeError::IdempotencyConflict { operation_id });
|
||||
}
|
||||
match existing.payload.get("state").and_then(Value::as_str) {
|
||||
Some("completed") => {
|
||||
let response = existing
|
||||
.payload
|
||||
.get("response")
|
||||
.cloned()
|
||||
.and_then(|value| serde_json::from_value(value).ok())
|
||||
.ok_or_else(|| RuntimeError::InvalidPreparedRequest {
|
||||
field: "idempotency_state".to_owned(),
|
||||
reason: "completed idempotency state has no valid response".to_owned(),
|
||||
})?;
|
||||
return Ok(IdempotencyAction::Replay(response));
|
||||
}
|
||||
Some("outcome_unknown") => {
|
||||
return Err(RuntimeError::IdempotencyOutcomeUnknown { operation_id });
|
||||
}
|
||||
Some("in_progress") if Instant::now() < deadline => {
|
||||
tokio::time::sleep(POLL_INTERVAL).await;
|
||||
existing = store
|
||||
.get_value(CacheScope::Coordination, key)
|
||||
.await
|
||||
.map_err(|_| RuntimeError::IdempotencyStoreUnavailable {
|
||||
operation_id: operation_id.clone(),
|
||||
})?
|
||||
.ok_or_else(|| RuntimeError::IdempotencyOutcomeUnknown {
|
||||
operation_id: operation_id.clone(),
|
||||
})?;
|
||||
}
|
||||
Some("in_progress") => {
|
||||
return Err(RuntimeError::IdempotencyInProgress { operation_id });
|
||||
}
|
||||
_ => {
|
||||
return Err(RuntimeError::InvalidPreparedRequest {
|
||||
field: "idempotency_state".to_owned(),
|
||||
reason: "unknown idempotency state".to_owned(),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn policy(operation: &RuntimeOperation) -> Option<&IdempotencyPolicy> {
|
||||
let is_mutating_rest =
|
||||
matches!(&operation.target, Target::Rest(target) if target.method != HttpMethod::Get);
|
||||
@@ -54,7 +256,7 @@ pub fn policy(operation: &RuntimeOperation) -> Option<&IdempotencyPolicy> {
|
||||
Some(policy)
|
||||
}
|
||||
|
||||
pub fn cache_key(
|
||||
fn cache_key(
|
||||
operation: &RuntimeOperation,
|
||||
idempotency_key: &str,
|
||||
request_context: Option<&RuntimeRequestContext>,
|
||||
@@ -76,6 +278,42 @@ pub fn cache_key(
|
||||
))
|
||||
}
|
||||
|
||||
fn request_fingerprint(input: &Value) -> Result<String, RuntimeError> {
|
||||
let canonical = canonical_json(input);
|
||||
let encoded =
|
||||
serde_json::to_vec(&canonical).map_err(|error| RuntimeError::InvalidPreparedRequest {
|
||||
field: "idempotency_fingerprint".to_owned(),
|
||||
reason: error.to_string(),
|
||||
})?;
|
||||
Ok(URL_SAFE_NO_PAD.encode(Sha256::digest(encoded)))
|
||||
}
|
||||
|
||||
fn canonical_json(value: &Value) -> Value {
|
||||
match value {
|
||||
Value::Object(object) => {
|
||||
let mut entries = object.iter().collect::<Vec<_>>();
|
||||
entries.sort_unstable_by_key(|(key, _)| *key);
|
||||
let mut canonical = Map::new();
|
||||
for (key, value) in entries {
|
||||
canonical.insert(key.clone(), canonical_json(value));
|
||||
}
|
||||
Value::Object(canonical)
|
||||
}
|
||||
Value::Array(values) => Value::Array(values.iter().map(canonical_json).collect()),
|
||||
_ => value.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
fn in_progress_value(fingerprint: &str) -> CoordinationStateValue {
|
||||
CoordinationStateValue {
|
||||
payload: json!({
|
||||
"state": "in_progress",
|
||||
"fingerprint": fingerprint,
|
||||
"owner": Uuid::now_v7().simple().to_string(),
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn key_from_policy(
|
||||
policy: &IdempotencyPolicy,
|
||||
input: &Value,
|
||||
|
||||
@@ -32,7 +32,8 @@ pub use executor_builder::{
|
||||
pub use limits::{RuntimeLimits, RuntimeLimitsConfigError};
|
||||
pub use model::{AdapterResponse, PreparedRequest, RuntimeOperation};
|
||||
pub use rate_limit::{
|
||||
RateLimitRejection, RequestRateLimitConfig, RequestRateLimitConfigError, RequestRateLimiter,
|
||||
RateLimitCheckError, RateLimitRejection, RequestRateLimitConfig, RequestRateLimitConfigError,
|
||||
RequestRateLimiter,
|
||||
};
|
||||
pub use request_context::{MeteringContext, ResponseCacheScope, RuntimeRequestContext};
|
||||
pub use secret_crypto::SecretCrypto;
|
||||
|
||||
@@ -3,16 +3,19 @@ use std::{env, num::ParseIntError};
|
||||
use thiserror::Error;
|
||||
|
||||
const DEFAULT_MAX_CONCURRENT_UNARY: usize = 64;
|
||||
const DEFAULT_MAX_CONCURRENT_SESSIONS: usize = 16;
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub struct RuntimeLimits {
|
||||
pub max_concurrent_unary: usize,
|
||||
pub max_concurrent_sessions: usize,
|
||||
}
|
||||
|
||||
impl Default for RuntimeLimits {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
max_concurrent_unary: DEFAULT_MAX_CONCURRENT_UNARY,
|
||||
max_concurrent_sessions: DEFAULT_MAX_CONCURRENT_SESSIONS,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -24,6 +27,10 @@ impl RuntimeLimits {
|
||||
"CRANK_RUNTIME_MAX_CONCURRENT_UNARY",
|
||||
DEFAULT_MAX_CONCURRENT_UNARY,
|
||||
)?,
|
||||
max_concurrent_sessions: parse_limit(
|
||||
"CRANK_RUNTIME_MAX_CONCURRENT_SESSIONS",
|
||||
DEFAULT_MAX_CONCURRENT_SESSIONS,
|
||||
)?,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -71,6 +78,7 @@ mod tests {
|
||||
let limits = RuntimeLimits::default();
|
||||
|
||||
assert!(limits.max_concurrent_unary > 0);
|
||||
assert!(limits.max_concurrent_sessions > 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -4,8 +4,9 @@ use std::{
|
||||
time::{Duration, Instant, SystemTime, UNIX_EPOCH},
|
||||
};
|
||||
|
||||
use crank_core::{RateLimitBucketState, RateLimitStateStore};
|
||||
use crank_core::{RateLimitDecision, RateLimitStateStore};
|
||||
use thiserror::Error;
|
||||
use tracing::warn;
|
||||
|
||||
const STALE_KEY_TTL: Duration = Duration::from_secs(300);
|
||||
const TOKEN_SCALE: u64 = 1_000_000;
|
||||
@@ -47,6 +48,12 @@ pub struct RateLimitRejection {
|
||||
pub retry_after_ms: u64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum RateLimitCheckError {
|
||||
Rejected(RateLimitRejection),
|
||||
StoreUnavailable,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct RequestRateLimiter {
|
||||
config: RequestRateLimitConfig,
|
||||
@@ -87,14 +94,24 @@ impl RequestRateLimiter {
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn check(&self, key: &str) -> Result<(), RateLimitRejection> {
|
||||
match &self.backend {
|
||||
RequestRateLimiterBackend::Local { .. } => self.check_local_at(key, Instant::now()),
|
||||
pub async fn check(&self, key: &str) -> Result<(), RateLimitCheckError> {
|
||||
let result = match &self.backend {
|
||||
RequestRateLimiterBackend::Local { .. } => self
|
||||
.check_local_at(key, Instant::now())
|
||||
.map_err(RateLimitCheckError::Rejected),
|
||||
RequestRateLimiterBackend::Shared { store } => {
|
||||
self.check_shared_at(store.as_ref(), key, now_unix_ms())
|
||||
.await
|
||||
}
|
||||
};
|
||||
if matches!(result, Err(RateLimitCheckError::Rejected(_))) {
|
||||
metrics::counter!(
|
||||
"crank_runtime_limit_rejections_total",
|
||||
"stage" => "rate_limit"
|
||||
)
|
||||
.increment(1);
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
fn check_local_at(&self, key: &str, now: Instant) -> Result<(), RateLimitRejection> {
|
||||
@@ -137,35 +154,35 @@ impl RequestRateLimiter {
|
||||
store: &dyn RateLimitStateStore,
|
||||
key: &str,
|
||||
now_unix_ms: i64,
|
||||
) -> Result<(), RateLimitRejection> {
|
||||
let burst_tokens = u64::from(self.config.burst) * TOKEN_SCALE;
|
||||
let refill_per_second = u64::from(self.config.requests_per_second) * TOKEN_SCALE;
|
||||
let mut state =
|
||||
store
|
||||
.get_bucket(key)
|
||||
.await
|
||||
.unwrap_or(None)
|
||||
.unwrap_or(RateLimitBucketState {
|
||||
tokens_micros: burst_tokens,
|
||||
last_refill_unix_ms: now_unix_ms,
|
||||
});
|
||||
|
||||
let elapsed_ms = (now_unix_ms - state.last_refill_unix_ms).max(0) as u64;
|
||||
let replenished =
|
||||
state.tokens_micros + (elapsed_ms.saturating_mul(refill_per_second) / 1000);
|
||||
state.tokens_micros = replenished.min(burst_tokens);
|
||||
state.last_refill_unix_ms = now_unix_ms;
|
||||
|
||||
if state.tokens_micros >= TOKEN_SCALE {
|
||||
state.tokens_micros -= TOKEN_SCALE;
|
||||
let _ = store.put_bucket(key, state, STALE_KEY_TTL).await;
|
||||
return Ok(());
|
||||
) -> Result<(), RateLimitCheckError> {
|
||||
let decision = match store
|
||||
.consume_token(
|
||||
key,
|
||||
u64::from(self.config.burst) * TOKEN_SCALE,
|
||||
u64::from(self.config.requests_per_second) * TOKEN_SCALE,
|
||||
now_unix_ms,
|
||||
STALE_KEY_TTL,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(decision) => decision,
|
||||
Err(_) => {
|
||||
warn!(
|
||||
name: "runtime.rate_limit.failed_closed",
|
||||
error_category = "coordination_store",
|
||||
"shared rate limiter failed closed"
|
||||
);
|
||||
return Err(RateLimitCheckError::StoreUnavailable);
|
||||
}
|
||||
};
|
||||
match decision {
|
||||
RateLimitDecision::Allowed => Ok(()),
|
||||
RateLimitDecision::Rejected { retry_after_ms } => {
|
||||
Err(RateLimitCheckError::Rejected(RateLimitRejection {
|
||||
retry_after_ms,
|
||||
}))
|
||||
}
|
||||
}
|
||||
|
||||
let missing_tokens = TOKEN_SCALE.saturating_sub(state.tokens_micros);
|
||||
let retry_after_ms = missing_tokens.div_ceil(refill_per_second).max(1);
|
||||
let _ = store.put_bucket(key, state, STALE_KEY_TTL).await;
|
||||
Err(RateLimitRejection { retry_after_ms })
|
||||
}
|
||||
}
|
||||
|
||||
@@ -185,6 +202,11 @@ mod tests {
|
||||
time::{Duration, Instant},
|
||||
};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use crank_core::{
|
||||
CacheStoreError, RateLimitBucketState, RateLimitDecision, RateLimitStateStore,
|
||||
};
|
||||
|
||||
use crate::InMemoryRateLimitStateStore;
|
||||
|
||||
use super::{RequestRateLimitConfig, RequestRateLimitConfigError, RequestRateLimiter};
|
||||
@@ -234,7 +256,12 @@ mod tests {
|
||||
assert!(limiter.check_shared_at_store("key", 0).await.is_ok());
|
||||
|
||||
let rejection = limiter.check_shared_at_store("key", 0).await.unwrap_err();
|
||||
assert_eq!(rejection.retry_after_ms, 1);
|
||||
assert_eq!(
|
||||
rejection,
|
||||
super::RateLimitCheckError::Rejected(super::RateLimitRejection {
|
||||
retry_after_ms: 500,
|
||||
})
|
||||
);
|
||||
|
||||
assert!(limiter.check_shared_at_store("key", 500).await.is_ok());
|
||||
}
|
||||
@@ -254,12 +281,120 @@ mod tests {
|
||||
assert!(second.check_shared_at_store("shared", 1000).await.is_ok());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn shared_limiter_does_not_refill_when_clock_moves_backwards() {
|
||||
let limiter = RequestRateLimiter::new_shared(
|
||||
RequestRateLimitConfig::new(1, 1).unwrap(),
|
||||
Arc::new(InMemoryRateLimitStateStore::default()),
|
||||
);
|
||||
|
||||
assert!(limiter.check_shared_at_store("clock", 1_000).await.is_ok());
|
||||
assert_eq!(
|
||||
limiter
|
||||
.check_shared_at_store("clock", 500)
|
||||
.await
|
||||
.unwrap_err(),
|
||||
super::RateLimitCheckError::Rejected(super::RateLimitRejection {
|
||||
retry_after_ms: 1_000,
|
||||
})
|
||||
);
|
||||
assert_eq!(
|
||||
limiter
|
||||
.check_shared_at_store("clock", 1_500)
|
||||
.await
|
||||
.unwrap_err(),
|
||||
super::RateLimitCheckError::Rejected(super::RateLimitRejection {
|
||||
retry_after_ms: 500,
|
||||
})
|
||||
);
|
||||
assert!(limiter.check_shared_at_store("clock", 2_000).await.is_ok());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn shared_limiter_consumes_burst_atomically_under_concurrency() {
|
||||
let limiter = RequestRateLimiter::new_shared(
|
||||
RequestRateLimitConfig::new(1, 8).unwrap(),
|
||||
Arc::new(InMemoryRateLimitStateStore::default()),
|
||||
);
|
||||
let attempts = (0..64).map(|_| limiter.check_shared_at_store("concurrent", 0));
|
||||
let results = futures_util::future::join_all(attempts).await;
|
||||
|
||||
assert_eq!(results.iter().filter(|result| result.is_ok()).count(), 8);
|
||||
assert!(
|
||||
results
|
||||
.iter()
|
||||
.filter_map(|result| result.as_ref().err())
|
||||
.all(|rejection| {
|
||||
*rejection
|
||||
== super::RateLimitCheckError::Rejected(super::RateLimitRejection {
|
||||
retry_after_ms: 1_000,
|
||||
})
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn shared_limiter_fails_closed_when_store_is_unavailable() {
|
||||
let limiter = RequestRateLimiter::new_shared(
|
||||
RequestRateLimitConfig::new(10, 10).unwrap(),
|
||||
Arc::new(UnavailableRateLimitStore),
|
||||
);
|
||||
|
||||
let error = limiter
|
||||
.check_shared_at_store("unavailable", 0)
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert_eq!(error, super::RateLimitCheckError::StoreUnavailable);
|
||||
}
|
||||
|
||||
struct UnavailableRateLimitStore;
|
||||
|
||||
#[async_trait]
|
||||
impl RateLimitStateStore for UnavailableRateLimitStore {
|
||||
async fn get_bucket(
|
||||
&self,
|
||||
_key: &str,
|
||||
) -> Result<Option<RateLimitBucketState>, CacheStoreError> {
|
||||
Err(unavailable())
|
||||
}
|
||||
|
||||
async fn put_bucket(
|
||||
&self,
|
||||
_key: &str,
|
||||
_value: RateLimitBucketState,
|
||||
_ttl: Duration,
|
||||
) -> Result<(), CacheStoreError> {
|
||||
Err(unavailable())
|
||||
}
|
||||
|
||||
async fn delete_bucket(&self, _key: &str) -> Result<(), CacheStoreError> {
|
||||
Err(unavailable())
|
||||
}
|
||||
|
||||
async fn consume_token(
|
||||
&self,
|
||||
_key: &str,
|
||||
_burst_tokens_micros: u64,
|
||||
_refill_per_second_micros: u64,
|
||||
_now_unix_ms: i64,
|
||||
_ttl: Duration,
|
||||
) -> Result<RateLimitDecision, CacheStoreError> {
|
||||
Err(unavailable())
|
||||
}
|
||||
}
|
||||
|
||||
fn unavailable() -> CacheStoreError {
|
||||
CacheStoreError::Unavailable {
|
||||
message: "test store is unavailable".to_owned(),
|
||||
}
|
||||
}
|
||||
|
||||
impl RequestRateLimiter {
|
||||
async fn check_shared_at_store(
|
||||
&self,
|
||||
key: &str,
|
||||
now_unix_ms: i64,
|
||||
) -> Result<(), super::RateLimitRejection> {
|
||||
) -> Result<(), super::RateLimitCheckError> {
|
||||
let super::RequestRateLimiterBackend::Shared { store } = &self.backend else {
|
||||
panic!("check_shared_at_store called for non-shared limiter");
|
||||
};
|
||||
|
||||
@@ -3,7 +3,7 @@ use std::collections::BTreeMap;
|
||||
use crank_core::Target;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{PreparedRequest, RuntimeError, RuntimeOperation, RuntimeRequestContext};
|
||||
use crate::{PreparedRequest, RuntimeError, RuntimeOperation};
|
||||
|
||||
impl PreparedRequest {
|
||||
pub fn from_mapping_output(mapped: &Value) -> Result<Self, RuntimeError> {
|
||||
@@ -28,7 +28,6 @@ impl PreparedRequest {
|
||||
pub(crate) fn adapter_prepared_request(
|
||||
operation: &RuntimeOperation,
|
||||
prepared_request: &PreparedRequest,
|
||||
request_context: Option<&RuntimeRequestContext>,
|
||||
timeout_ms: u64,
|
||||
) -> PreparedRequest {
|
||||
let static_headers = match &operation.target {
|
||||
@@ -40,7 +39,6 @@ pub(crate) fn adapter_prepared_request(
|
||||
static_headers,
|
||||
&operation.execution_config.headers,
|
||||
&prepared_request.headers,
|
||||
&runtime_context_headers(request_context),
|
||||
);
|
||||
prepared_request.timeout_ms = timeout_ms;
|
||||
prepared_request
|
||||
@@ -50,23 +48,13 @@ fn merge_headers(
|
||||
static_headers: &BTreeMap<String, String>,
|
||||
execution_headers: &BTreeMap<String, String>,
|
||||
request_headers: &BTreeMap<String, String>,
|
||||
context_headers: &BTreeMap<String, String>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let mut headers = static_headers.clone();
|
||||
headers.extend(execution_headers.clone());
|
||||
headers.extend(request_headers.clone());
|
||||
headers.extend(context_headers.clone());
|
||||
headers
|
||||
}
|
||||
|
||||
fn runtime_context_headers(
|
||||
request_context: Option<&RuntimeRequestContext>,
|
||||
) -> BTreeMap<String, String> {
|
||||
request_context
|
||||
.map(RuntimeRequestContext::outbound_headers)
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
fn read_string_map(
|
||||
value: Option<&Value>,
|
||||
field_name: &str,
|
||||
|
||||
@@ -2,4 +2,6 @@ mod integration {
|
||||
mod confirmation;
|
||||
mod idempotency;
|
||||
mod no_input_get;
|
||||
mod stages;
|
||||
mod valkey;
|
||||
}
|
||||
|
||||
@@ -16,6 +16,7 @@ use crank_runtime::{
|
||||
InMemoryCoordinationStateStore, RuntimeError, RuntimeExecutorBuilder, RuntimeRequestContext,
|
||||
};
|
||||
use crank_schema::{Schema, SchemaKind};
|
||||
use futures_util::future::join_all;
|
||||
use serde_json::json;
|
||||
use time::OffsetDateTime;
|
||||
|
||||
@@ -89,6 +90,62 @@ async fn destructive_operation_requires_single_use_confirmation() {
|
||||
assert_eq!(call_count.load(Ordering::SeqCst), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn confirmation_token_allows_only_one_concurrent_execution() {
|
||||
let call_count = Arc::new(AtomicUsize::new(0));
|
||||
let executor = RuntimeExecutorBuilder::new()
|
||||
.register_adapter(Arc::new(CountingAdapter {
|
||||
call_count: Arc::clone(&call_count),
|
||||
}))
|
||||
.with_coordination_store(Arc::new(InMemoryCoordinationStateStore::default()))
|
||||
.build();
|
||||
let operation: crank_runtime::RuntimeOperation = destructive_delete_operation().into();
|
||||
let context = RuntimeRequestContext::from_request_id("req_confirm_concurrent")
|
||||
.with_response_cache_scope("workspace_1", "agent_1");
|
||||
let first = executor
|
||||
.execute_with_context(
|
||||
&operation,
|
||||
&json!({ "order_id": "ord_123" }),
|
||||
Some(&context),
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
let RuntimeError::ConfirmationRequired {
|
||||
confirmation_token, ..
|
||||
} = first
|
||||
else {
|
||||
panic!("expected confirmation token")
|
||||
};
|
||||
|
||||
let attempts = (0..16).map(|_| {
|
||||
let executor = executor.clone();
|
||||
let operation = operation.clone();
|
||||
let context = context
|
||||
.clone()
|
||||
.with_confirmation_token(confirmation_token.clone());
|
||||
async move {
|
||||
executor
|
||||
.execute_with_context(
|
||||
&operation,
|
||||
&json!({ "order_id": "ord_123" }),
|
||||
Some(&context),
|
||||
)
|
||||
.await
|
||||
}
|
||||
});
|
||||
let results = join_all(attempts).await;
|
||||
let successful = results.iter().filter(|result| result.is_ok()).count();
|
||||
|
||||
assert_eq!(successful, 1);
|
||||
assert_eq!(call_count.load(Ordering::SeqCst), 1);
|
||||
assert!(
|
||||
results
|
||||
.iter()
|
||||
.filter_map(|result| result.as_ref().err())
|
||||
.all(|error| matches!(error, RuntimeError::InvalidConfirmationToken { .. }))
|
||||
);
|
||||
}
|
||||
|
||||
struct CountingAdapter {
|
||||
call_count: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
@@ -12,7 +12,10 @@ use crank_core::{
|
||||
ToolDescription,
|
||||
};
|
||||
use crank_mapping::{MappingRule, MappingSet};
|
||||
use crank_runtime::{InMemoryResponseCacheStore, RuntimeExecutorBuilder};
|
||||
use crank_runtime::{
|
||||
InMemoryCoordinationStateStore, InMemoryResponseCacheStore, RuntimeError,
|
||||
RuntimeExecutorBuilder,
|
||||
};
|
||||
use crank_schema::{Schema, SchemaKind};
|
||||
use serde_json::json;
|
||||
use time::OffsetDateTime;
|
||||
@@ -23,8 +26,10 @@ async fn replays_mutation_result_for_same_idempotency_key() {
|
||||
let executor = RuntimeExecutorBuilder::new()
|
||||
.register_adapter(Arc::new(CountingAdapter {
|
||||
call_count: Arc::clone(&call_count),
|
||||
release: None,
|
||||
}))
|
||||
.with_response_cache(Arc::new(InMemoryResponseCacheStore::default()))
|
||||
.with_coordination_store(Arc::new(InMemoryCoordinationStateStore::default()))
|
||||
.build();
|
||||
let operation = idempotent_post_operation().into();
|
||||
let context = crank_runtime::RuntimeRequestContext::from_request_id("req_1")
|
||||
@@ -51,8 +56,10 @@ async fn required_idempotency_rejects_missing_key_before_adapter_call() {
|
||||
let executor = RuntimeExecutorBuilder::new()
|
||||
.register_adapter(Arc::new(CountingAdapter {
|
||||
call_count: Arc::clone(&call_count),
|
||||
release: None,
|
||||
}))
|
||||
.with_response_cache(Arc::new(InMemoryResponseCacheStore::default()))
|
||||
.with_coordination_store(Arc::new(InMemoryCoordinationStateStore::default()))
|
||||
.build();
|
||||
let mut operation: crank_runtime::RuntimeOperation = idempotent_post_operation().into();
|
||||
let policy = operation.execution_config.idempotency.as_mut().unwrap();
|
||||
@@ -71,8 +78,189 @@ async fn required_idempotency_rejects_missing_key_before_adapter_call() {
|
||||
assert_eq!(call_count.load(Ordering::SeqCst), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn required_idempotency_fails_closed_without_coordination_store() {
|
||||
let call_count = Arc::new(AtomicUsize::new(0));
|
||||
let executor = RuntimeExecutorBuilder::new()
|
||||
.register_adapter(Arc::new(CountingAdapter {
|
||||
call_count: Arc::clone(&call_count),
|
||||
release: None,
|
||||
}))
|
||||
.build();
|
||||
let operation = idempotent_post_operation().into();
|
||||
let context = crank_runtime::RuntimeRequestContext::from_request_id("req_no_store")
|
||||
.with_response_cache_scope("workspace_1", "agent_1");
|
||||
|
||||
let error = executor
|
||||
.execute_with_context(
|
||||
&operation,
|
||||
&json!({ "request_id": "must-not-run" }),
|
||||
Some(&context),
|
||||
)
|
||||
.await
|
||||
.expect_err("required idempotency must not execute without an atomic store");
|
||||
|
||||
assert!(matches!(
|
||||
error,
|
||||
RuntimeError::IdempotencyStoreUnavailable { .. }
|
||||
));
|
||||
assert_eq!(call_count.load(Ordering::SeqCst), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn concurrent_calls_with_same_key_execute_adapter_once() {
|
||||
let call_count = Arc::new(AtomicUsize::new(0));
|
||||
let release = Arc::new(tokio::sync::Notify::new());
|
||||
let executor = RuntimeExecutorBuilder::new()
|
||||
.register_adapter(Arc::new(CountingAdapter {
|
||||
call_count: Arc::clone(&call_count),
|
||||
release: Some(Arc::clone(&release)),
|
||||
}))
|
||||
.with_coordination_store(Arc::new(InMemoryCoordinationStateStore::default()))
|
||||
.build();
|
||||
let operation: crank_runtime::RuntimeOperation = idempotent_post_operation().into();
|
||||
let context = crank_runtime::RuntimeRequestContext::from_request_id("req_concurrent")
|
||||
.with_response_cache_scope("workspace_1", "agent_1");
|
||||
let input = json!({ "request_id": "order-concurrent" });
|
||||
|
||||
let first_executor = executor.clone();
|
||||
let first_operation = operation.clone();
|
||||
let first_context = context.clone();
|
||||
let first_input = input.clone();
|
||||
let first = tokio::spawn(async move {
|
||||
first_executor
|
||||
.execute_with_context(&first_operation, &first_input, Some(&first_context))
|
||||
.await
|
||||
});
|
||||
wait_for_call_count(&call_count, 1).await;
|
||||
|
||||
let second_executor = executor.clone();
|
||||
let second_operation = operation.clone();
|
||||
let second_context = context.clone();
|
||||
let second_input = input.clone();
|
||||
let second = tokio::spawn(async move {
|
||||
second_executor
|
||||
.execute_with_context(&second_operation, &second_input, Some(&second_context))
|
||||
.await
|
||||
});
|
||||
for _ in 0..100 {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
assert_eq!(call_count.load(Ordering::SeqCst), 1);
|
||||
|
||||
release.notify_waiters();
|
||||
let first_result = first.await.unwrap().unwrap();
|
||||
let second_result = second.await.unwrap().unwrap();
|
||||
assert_eq!(first_result, second_result);
|
||||
assert_eq!(call_count.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn same_key_with_different_input_is_rejected() {
|
||||
let call_count = Arc::new(AtomicUsize::new(0));
|
||||
let executor = RuntimeExecutorBuilder::new()
|
||||
.register_adapter(Arc::new(CountingAdapter {
|
||||
call_count: Arc::clone(&call_count),
|
||||
release: None,
|
||||
}))
|
||||
.with_coordination_store(Arc::new(InMemoryCoordinationStateStore::default()))
|
||||
.build();
|
||||
let operation = idempotent_post_operation().into();
|
||||
let context = crank_runtime::RuntimeRequestContext::from_request_id("req_conflict")
|
||||
.with_response_cache_scope("workspace_1", "agent_1");
|
||||
|
||||
executor
|
||||
.execute_with_context(
|
||||
&operation,
|
||||
&json!({ "request_id": "stable-key", "amount": 10 }),
|
||||
Some(&context),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let error = executor
|
||||
.execute_with_context(
|
||||
&operation,
|
||||
&json!({ "request_id": "stable-key", "amount": 20 }),
|
||||
Some(&context),
|
||||
)
|
||||
.await
|
||||
.expect_err("same key must not accept a different request fingerprint");
|
||||
|
||||
assert!(matches!(error, RuntimeError::IdempotencyConflict { .. }));
|
||||
assert_eq!(call_count.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn uncertain_adapter_failure_blocks_automatic_retry() {
|
||||
let call_count = Arc::new(AtomicUsize::new(0));
|
||||
let executor = RuntimeExecutorBuilder::new()
|
||||
.register_adapter(Arc::new(FailingAdapter {
|
||||
call_count: Arc::clone(&call_count),
|
||||
}))
|
||||
.with_coordination_store(Arc::new(InMemoryCoordinationStateStore::default()))
|
||||
.build();
|
||||
let operation = idempotent_post_operation().into();
|
||||
let context = crank_runtime::RuntimeRequestContext::from_request_id("req_unknown")
|
||||
.with_response_cache_scope("workspace_1", "agent_1");
|
||||
let input = json!({ "request_id": "uncertain-outcome" });
|
||||
|
||||
let first = executor
|
||||
.execute_with_context(&operation, &input, Some(&context))
|
||||
.await
|
||||
.expect_err("adapter failure must be returned");
|
||||
assert!(matches!(first, RuntimeError::ProtocolAdapter(_)));
|
||||
|
||||
let retry = executor
|
||||
.execute_with_context(&operation, &input, Some(&context))
|
||||
.await
|
||||
.expect_err("an uncertain external outcome must not be retried automatically");
|
||||
assert!(matches!(
|
||||
retry,
|
||||
RuntimeError::IdempotencyOutcomeUnknown { .. }
|
||||
));
|
||||
assert_eq!(call_count.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
|
||||
async fn wait_for_call_count(call_count: &AtomicUsize, expected: usize) {
|
||||
for _ in 0..1_000 {
|
||||
if call_count.load(Ordering::SeqCst) >= expected {
|
||||
return;
|
||||
}
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
panic!("adapter did not receive {expected} call(s)");
|
||||
}
|
||||
|
||||
struct CountingAdapter {
|
||||
call_count: Arc<AtomicUsize>,
|
||||
release: Option<Arc<tokio::sync::Notify>>,
|
||||
}
|
||||
|
||||
struct FailingAdapter {
|
||||
call_count: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ProtocolAdapter for FailingAdapter {
|
||||
fn protocol(&self) -> Protocol {
|
||||
Protocol::Rest
|
||||
}
|
||||
|
||||
fn supports_mode(&self, mode: ExecutionMode) -> bool {
|
||||
mode == ExecutionMode::Unary
|
||||
}
|
||||
|
||||
async fn invoke_unary(
|
||||
&self,
|
||||
_target: &Target,
|
||||
_prepared: &crank_core::PreparedRequest,
|
||||
_context: &RuntimeRequestContext,
|
||||
) -> Result<AdapterResponse, ProtocolAdapterError> {
|
||||
self.call_count.fetch_add(1, Ordering::SeqCst);
|
||||
Err(ProtocolAdapterError::Message(
|
||||
"upstream outcome is unknown".to_owned(),
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -91,11 +279,11 @@ impl ProtocolAdapter for CountingAdapter {
|
||||
prepared: &crank_core::PreparedRequest,
|
||||
_context: &RuntimeRequestContext,
|
||||
) -> Result<AdapterResponse, ProtocolAdapterError> {
|
||||
assert_eq!(
|
||||
prepared.headers.get("Idempotency-Key").map(String::as_str),
|
||||
Some("order-123")
|
||||
);
|
||||
assert!(prepared.headers.contains_key("Idempotency-Key"));
|
||||
let call_number = self.call_count.fetch_add(1, Ordering::SeqCst) + 1;
|
||||
if let Some(release) = &self.release {
|
||||
release.notified().await;
|
||||
}
|
||||
Ok(AdapterResponse {
|
||||
status_code: 201,
|
||||
headers: BTreeMap::new(),
|
||||
|
||||
@@ -0,0 +1,426 @@
|
||||
use std::{
|
||||
collections::BTreeMap,
|
||||
sync::{Arc, Mutex},
|
||||
};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use crank_core::{
|
||||
AdapterResponse, ConfirmationPolicy, ExecutionConfig, ExecutionMode, HttpMethod,
|
||||
IdempotencyMode, IdempotencyPolicy, Operation, OperationId, OperationSafetyClass,
|
||||
OperationSafetyPolicy, OperationSecurityLevel, OperationStatus, Protocol, ProtocolAdapter,
|
||||
ProtocolAdapterError, RestTarget, Target, ToolDescription,
|
||||
};
|
||||
use crank_mapping::{MappingRule, MappingSet};
|
||||
use crank_runtime::{
|
||||
InMemoryCoordinationStateStore, RuntimeError, RuntimeExecutorBuilder, RuntimeRequestContext,
|
||||
};
|
||||
use crank_schema::{Schema, SchemaKind};
|
||||
use crank_trace::{Stage, StageOutcome};
|
||||
use serde_json::json;
|
||||
use time::OffsetDateTime;
|
||||
use tracing::{Id, Instrument, Subscriber, field::Visit, instrument::WithSubscriber};
|
||||
use tracing_subscriber::{Layer, layer::SubscriberExt, registry::LookupSpan};
|
||||
|
||||
#[tokio::test]
|
||||
async fn successful_execution_has_real_stages_and_omits_inapplicable_ones() {
|
||||
let capture = TraceCapture::default();
|
||||
let subscriber = tracing_subscriber::registry().with(capture.clone());
|
||||
let executor = RuntimeExecutorBuilder::new()
|
||||
.register_adapter(Arc::new(SuccessAdapter))
|
||||
.build();
|
||||
let operation = operation().into();
|
||||
let context = RuntimeRequestContext::from_request_id("req_stage_test");
|
||||
let result = async {
|
||||
let root = tracing::info_span!(target: "crank::trace", "mcp.request");
|
||||
executor
|
||||
.execute_with_context(
|
||||
&operation,
|
||||
&json!({"name": "canary-secret"}),
|
||||
Some(&context),
|
||||
)
|
||||
.instrument(root)
|
||||
.await
|
||||
}
|
||||
.with_subscriber(subscriber)
|
||||
.await;
|
||||
|
||||
assert_eq!(result.unwrap(), json!({"accepted": true}));
|
||||
let spans = capture.snapshot();
|
||||
assert_stage(&spans, "runtime.execute", "success");
|
||||
assert_stage(&spans, "runtime.arguments.map", "success");
|
||||
assert_stage(&spans, "upstream.http", "success");
|
||||
assert_stage(&spans, "runtime.response.transform", "success");
|
||||
assert!(!spans.iter().any(|span| span.name == "approval.check"));
|
||||
assert!(!spans.iter().any(|span| span.name == "runtime.idempotency"));
|
||||
assert!(
|
||||
spans
|
||||
.iter()
|
||||
.flat_map(|span| span.fields.values())
|
||||
.all(|value| !value.contains("canary-secret"))
|
||||
);
|
||||
|
||||
let runtime = spans
|
||||
.iter()
|
||||
.find(|span| span.name == "runtime.execute")
|
||||
.expect("runtime span");
|
||||
assert_eq!(runtime.parent_name, Some("mcp.request"));
|
||||
for child in [
|
||||
"runtime.arguments.map",
|
||||
"upstream.http",
|
||||
"runtime.response.transform",
|
||||
] {
|
||||
assert_eq!(
|
||||
spans
|
||||
.iter()
|
||||
.find(|span| span.name == child)
|
||||
.and_then(|span| span.parent_name),
|
||||
Some("runtime.execute"),
|
||||
"{child} must be a runtime child"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn failed_mapping_records_closed_category_and_stops_later_stages() {
|
||||
let capture = TraceCapture::default();
|
||||
let subscriber = tracing_subscriber::registry().with(capture.clone());
|
||||
let executor = RuntimeExecutorBuilder::new()
|
||||
.register_adapter(Arc::new(SuccessAdapter))
|
||||
.build();
|
||||
let operation = operation().into();
|
||||
|
||||
let result = async { executor.execute(&operation, &json!({})).await }
|
||||
.with_subscriber(subscriber)
|
||||
.await;
|
||||
|
||||
assert!(result.is_err());
|
||||
let spans = capture.snapshot();
|
||||
let runtime = spans
|
||||
.iter()
|
||||
.find(|span| span.name == "runtime.execute")
|
||||
.expect("runtime span");
|
||||
assert_eq!(runtime.fields["outcome"], "error");
|
||||
assert_eq!(runtime.fields["error.category"], "schema");
|
||||
assert_stage(&spans, "runtime.arguments.map", "error");
|
||||
assert!(!spans.iter().any(|span| span.name == "upstream.http"));
|
||||
assert!(
|
||||
!spans
|
||||
.iter()
|
||||
.any(|span| span.name == "runtime.response.transform")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn approval_stage_is_present_only_when_confirmation_is_required() {
|
||||
let capture = TraceCapture::default();
|
||||
let subscriber = tracing_subscriber::registry().with(capture.clone());
|
||||
let executor = RuntimeExecutorBuilder::new()
|
||||
.register_adapter(Arc::new(SuccessAdapter))
|
||||
.with_coordination_store(Arc::new(InMemoryCoordinationStateStore::default()))
|
||||
.build();
|
||||
let mut source = operation();
|
||||
source.execution_config.safety = Some(OperationSafetyPolicy {
|
||||
class: OperationSafetyClass::Destructive,
|
||||
confirmation: Some(ConfirmationPolicy { ttl_ms: 60_000 }),
|
||||
});
|
||||
let operation = source.into();
|
||||
let context = RuntimeRequestContext::from_request_id("req_approval_stage")
|
||||
.with_response_cache_scope("workspace", "agent");
|
||||
|
||||
let result = async {
|
||||
executor
|
||||
.execute_with_context(
|
||||
&operation,
|
||||
&json!({"name": "requires-confirmation"}),
|
||||
Some(&context),
|
||||
)
|
||||
.await
|
||||
}
|
||||
.with_subscriber(subscriber)
|
||||
.await;
|
||||
|
||||
assert!(matches!(
|
||||
result,
|
||||
Err(RuntimeError::ConfirmationRequired { .. })
|
||||
));
|
||||
let spans = capture.snapshot();
|
||||
assert_stage(&spans, "approval.check", "required");
|
||||
assert!(!spans.iter().any(|span| span.name == "upstream.http"));
|
||||
assert!(!spans.iter().any(|span| span.name == "runtime.idempotency"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn idempotency_stage_distinguishes_execution_from_replay() {
|
||||
let capture = TraceCapture::default();
|
||||
let subscriber = tracing_subscriber::registry().with(capture.clone());
|
||||
let executor = RuntimeExecutorBuilder::new()
|
||||
.register_adapter(Arc::new(SuccessAdapter))
|
||||
.with_coordination_store(Arc::new(InMemoryCoordinationStateStore::default()))
|
||||
.build();
|
||||
let mut source = operation();
|
||||
source.execution_config.idempotency = Some(IdempotencyPolicy {
|
||||
mode: IdempotencyMode::Required,
|
||||
ttl_ms: 60_000,
|
||||
input_field: Some("name".to_owned()),
|
||||
header_name: Some("Idempotency-Key".to_owned()),
|
||||
});
|
||||
let operation = source.into();
|
||||
let context = RuntimeRequestContext::from_request_id("req_idempotency_stage")
|
||||
.with_response_cache_scope("workspace", "agent");
|
||||
|
||||
let (first, replay) = async {
|
||||
let first = executor
|
||||
.execute_with_context(&operation, &json!({"name": "stable-key"}), Some(&context))
|
||||
.await;
|
||||
let replay = executor
|
||||
.execute_with_context(&operation, &json!({"name": "stable-key"}), Some(&context))
|
||||
.await;
|
||||
(first, replay)
|
||||
}
|
||||
.with_subscriber(subscriber)
|
||||
.await;
|
||||
|
||||
assert!(first.is_ok());
|
||||
assert!(replay.is_ok());
|
||||
let spans = capture.snapshot();
|
||||
let idempotency_outcomes = spans
|
||||
.iter()
|
||||
.filter(|span| span.name == "runtime.idempotency")
|
||||
.map(|span| span.fields["outcome"].as_str())
|
||||
.collect::<Vec<_>>();
|
||||
assert!(idempotency_outcomes.contains(&"execute"));
|
||||
assert!(idempotency_outcomes.contains(&"replay"));
|
||||
assert_eq!(
|
||||
spans
|
||||
.iter()
|
||||
.filter(|span| span.name == "upstream.http")
|
||||
.count(),
|
||||
1,
|
||||
"replay must not pretend to call upstream"
|
||||
);
|
||||
}
|
||||
|
||||
fn assert_stage(spans: &[CapturedSpan], name: &str, outcome: &str) {
|
||||
let span = spans
|
||||
.iter()
|
||||
.find(|span| span.name == name)
|
||||
.unwrap_or_else(|| panic!("missing stage {name}"));
|
||||
assert_eq!(span.fields["outcome"], outcome);
|
||||
}
|
||||
|
||||
struct SuccessAdapter;
|
||||
|
||||
#[async_trait]
|
||||
impl ProtocolAdapter for SuccessAdapter {
|
||||
fn protocol(&self) -> Protocol {
|
||||
Protocol::Rest
|
||||
}
|
||||
|
||||
fn supports_mode(&self, mode: ExecutionMode) -> bool {
|
||||
mode == ExecutionMode::Unary
|
||||
}
|
||||
|
||||
async fn invoke_unary(
|
||||
&self,
|
||||
_target: &Target,
|
||||
_prepared: &crank_core::PreparedRequest,
|
||||
_context: &crank_core::RuntimeRequestContext,
|
||||
) -> Result<AdapterResponse, ProtocolAdapterError> {
|
||||
let span = Stage::UpstreamHttp.span();
|
||||
let response = Ok(AdapterResponse {
|
||||
status_code: 200,
|
||||
headers: BTreeMap::new(),
|
||||
body: json!({"accepted": true}),
|
||||
data: json!({"accepted": true}),
|
||||
});
|
||||
StageOutcome::Success.record(&span);
|
||||
response
|
||||
}
|
||||
}
|
||||
|
||||
fn operation() -> Operation<Schema, MappingSet> {
|
||||
Operation {
|
||||
id: OperationId::new("op_stage_test"),
|
||||
name: "stage_test".to_owned(),
|
||||
display_name: "Stage test".to_owned(),
|
||||
category: "test".to_owned(),
|
||||
protocol: Protocol::Rest,
|
||||
security_level: OperationSecurityLevel::Standard,
|
||||
status: OperationStatus::Published,
|
||||
version: 1,
|
||||
target: Target::Rest(RestTarget {
|
||||
base_url: "https://example.invalid".to_owned(),
|
||||
method: HttpMethod::Post,
|
||||
path_template: "/test".to_owned(),
|
||||
static_headers: BTreeMap::new(),
|
||||
}),
|
||||
input_schema: object_schema(BTreeMap::from([("name".to_owned(), string_schema())])),
|
||||
output_schema: object_schema(BTreeMap::from([("accepted".to_owned(), bool_schema())])),
|
||||
input_mapping: MappingSet {
|
||||
rules: vec![MappingRule {
|
||||
source: "$.mcp.name".to_owned(),
|
||||
target: "$.request.body.name".to_owned(),
|
||||
required: true,
|
||||
default_value: None,
|
||||
transform: None,
|
||||
condition: None,
|
||||
notes: None,
|
||||
}],
|
||||
},
|
||||
output_mapping: MappingSet {
|
||||
rules: vec![MappingRule {
|
||||
source: "$.response.body.accepted".to_owned(),
|
||||
target: "$.output.accepted".to_owned(),
|
||||
required: true,
|
||||
default_value: None,
|
||||
transform: None,
|
||||
condition: None,
|
||||
notes: None,
|
||||
}],
|
||||
},
|
||||
execution_config: ExecutionConfig {
|
||||
timeout_ms: 1_000,
|
||||
retry_policy: None,
|
||||
response_cache: None,
|
||||
idempotency: None,
|
||||
safety: None,
|
||||
approval_policy: None,
|
||||
auth_profile_ref: None,
|
||||
headers: BTreeMap::new(),
|
||||
},
|
||||
tool_description: ToolDescription {
|
||||
title: "Stage test".to_owned(),
|
||||
description: "Tests trace stages.".to_owned(),
|
||||
tags: Vec::new(),
|
||||
examples: Vec::new(),
|
||||
},
|
||||
samples: None,
|
||||
generated_draft: None,
|
||||
config_export: None,
|
||||
wizard_state: None,
|
||||
created_at: OffsetDateTime::UNIX_EPOCH,
|
||||
updated_at: OffsetDateTime::UNIX_EPOCH,
|
||||
published_at: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn object_schema(fields: BTreeMap<String, Schema>) -> Schema {
|
||||
Schema {
|
||||
kind: SchemaKind::Object,
|
||||
description: None,
|
||||
required: true,
|
||||
nullable: false,
|
||||
default_value: None,
|
||||
fields,
|
||||
items: None,
|
||||
enum_values: Vec::new(),
|
||||
variants: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn string_schema() -> Schema {
|
||||
Schema {
|
||||
kind: SchemaKind::String,
|
||||
description: None,
|
||||
required: true,
|
||||
nullable: false,
|
||||
default_value: None,
|
||||
fields: BTreeMap::new(),
|
||||
items: None,
|
||||
enum_values: Vec::new(),
|
||||
variants: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn bool_schema() -> Schema {
|
||||
Schema {
|
||||
kind: SchemaKind::Boolean,
|
||||
..string_schema()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
struct TraceCapture {
|
||||
spans: Arc<Mutex<Vec<CapturedSpan>>>,
|
||||
}
|
||||
|
||||
impl TraceCapture {
|
||||
fn snapshot(&self) -> Vec<CapturedSpan> {
|
||||
self.spans.lock().expect("span lock").clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
struct CapturedSpan {
|
||||
name: &'static str,
|
||||
parent_name: Option<&'static str>,
|
||||
fields: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
impl<S> Layer<S> for TraceCapture
|
||||
where
|
||||
S: Subscriber + for<'lookup> LookupSpan<'lookup>,
|
||||
{
|
||||
fn on_new_span(
|
||||
&self,
|
||||
attributes: &tracing::span::Attributes<'_>,
|
||||
id: &Id,
|
||||
context: tracing_subscriber::layer::Context<'_, S>,
|
||||
) {
|
||||
let parent = attributes
|
||||
.parent()
|
||||
.and_then(|parent| context.span(parent))
|
||||
.or_else(|| {
|
||||
attributes
|
||||
.is_contextual()
|
||||
.then(|| context.lookup_current())
|
||||
.flatten()
|
||||
});
|
||||
let mut visitor = FieldVisitor::default();
|
||||
attributes.record(&mut visitor);
|
||||
let mut spans = self.spans.lock().expect("span lock");
|
||||
let index = spans.len();
|
||||
spans.push(CapturedSpan {
|
||||
name: attributes.metadata().name(),
|
||||
parent_name: parent.map(|span| span.metadata().name()),
|
||||
fields: visitor.fields,
|
||||
});
|
||||
context
|
||||
.span(id)
|
||||
.expect("span exists")
|
||||
.extensions_mut()
|
||||
.insert(index);
|
||||
}
|
||||
|
||||
fn on_record(
|
||||
&self,
|
||||
id: &Id,
|
||||
values: &tracing::span::Record<'_>,
|
||||
context: tracing_subscriber::layer::Context<'_, S>,
|
||||
) {
|
||||
let mut visitor = FieldVisitor::default();
|
||||
values.record(&mut visitor);
|
||||
let span = context.span(id).expect("span exists");
|
||||
let index = *span.extensions().get::<usize>().expect("capture index");
|
||||
self.spans.lock().expect("span lock")[index]
|
||||
.fields
|
||||
.extend(visitor.fields);
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct FieldVisitor {
|
||||
fields: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
impl Visit for FieldVisitor {
|
||||
fn record_str(&mut self, field: &tracing::field::Field, value: &str) {
|
||||
self.fields
|
||||
.insert(field.name().to_owned(), value.to_owned());
|
||||
}
|
||||
|
||||
fn record_debug(&mut self, field: &tracing::field::Field, value: &dyn std::fmt::Debug) {
|
||||
self.fields
|
||||
.insert(field.name().to_owned(), format!("{value:?}"));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,167 @@
|
||||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use crank_core::{
|
||||
CacheBackend, CacheScope, CoordinationStateReservation, CoordinationStateStore,
|
||||
CoordinationStateValue, RateLimitDecision, RateLimitStateStore,
|
||||
};
|
||||
use crank_runtime::RedisCacheStore;
|
||||
use futures_util::future::join_all;
|
||||
use serde_json::json;
|
||||
use testcontainers::{
|
||||
GenericImage,
|
||||
core::{IntoContainerPort, WaitFor},
|
||||
runners::AsyncRunner,
|
||||
};
|
||||
|
||||
#[tokio::test]
|
||||
async fn valkey_coordination_and_rate_limit_operations_are_atomic() {
|
||||
let container = GenericImage::new("valkey/valkey", "8-alpine")
|
||||
.with_exposed_port(6379.tcp())
|
||||
.with_wait_for(WaitFor::message_on_stdout("Ready to accept connections"))
|
||||
.start()
|
||||
.await
|
||||
.expect("Valkey test container must start");
|
||||
let port = container
|
||||
.get_host_port_ipv4(6379.tcp())
|
||||
.await
|
||||
.expect("Valkey port must be mapped");
|
||||
let store = Arc::new(
|
||||
RedisCacheStore::connect(CacheBackend::Valkey, &format!("redis://127.0.0.1:{port}/0"))
|
||||
.await
|
||||
.expect("runtime store must connect to Valkey"),
|
||||
);
|
||||
|
||||
verify_atomic_coordination(store.as_ref()).await;
|
||||
verify_atomic_rate_limit(store).await;
|
||||
}
|
||||
|
||||
async fn verify_atomic_coordination(store: &RedisCacheStore) {
|
||||
let pending = CoordinationStateValue {
|
||||
payload: json!({ "state": "pending" }),
|
||||
};
|
||||
let attempts = (0..32).map(|_| {
|
||||
store.reserve_value(
|
||||
CacheScope::Coordination,
|
||||
"valkey-reservation",
|
||||
pending.clone(),
|
||||
Duration::from_secs(30),
|
||||
)
|
||||
});
|
||||
let results = join_all(attempts).await;
|
||||
assert_eq!(
|
||||
results
|
||||
.iter()
|
||||
.filter(|result| matches!(result, Ok(CoordinationStateReservation::Reserved)))
|
||||
.count(),
|
||||
1
|
||||
);
|
||||
|
||||
let completed = CoordinationStateValue {
|
||||
payload: json!({ "state": "completed" }),
|
||||
};
|
||||
assert!(
|
||||
store
|
||||
.compare_and_set_value(
|
||||
CacheScope::Coordination,
|
||||
"valkey-reservation",
|
||||
&pending,
|
||||
completed.clone(),
|
||||
Duration::from_secs(30),
|
||||
)
|
||||
.await
|
||||
.unwrap()
|
||||
);
|
||||
assert_eq!(
|
||||
store
|
||||
.take_value(CacheScope::Coordination, "valkey-reservation")
|
||||
.await
|
||||
.unwrap(),
|
||||
Some(completed)
|
||||
);
|
||||
assert_eq!(
|
||||
store
|
||||
.take_value(CacheScope::Coordination, "valkey-reservation")
|
||||
.await
|
||||
.unwrap(),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
async fn verify_atomic_rate_limit(store: Arc<RedisCacheStore>) {
|
||||
let attempts = (0..64).map(|_| {
|
||||
let store = Arc::clone(&store);
|
||||
async move {
|
||||
store
|
||||
.consume_token(
|
||||
"valkey-burst",
|
||||
8_000_000,
|
||||
1_000_000,
|
||||
0,
|
||||
Duration::from_secs(30),
|
||||
)
|
||||
.await
|
||||
}
|
||||
});
|
||||
let results = join_all(attempts).await;
|
||||
assert_eq!(
|
||||
results
|
||||
.iter()
|
||||
.filter(|result| matches!(result, Ok(RateLimitDecision::Allowed)))
|
||||
.count(),
|
||||
8
|
||||
);
|
||||
assert!(
|
||||
results
|
||||
.iter()
|
||||
.filter_map(|result| result.as_ref().ok())
|
||||
.all(|decision| matches!(
|
||||
decision,
|
||||
RateLimitDecision::Allowed
|
||||
| RateLimitDecision::Rejected {
|
||||
retry_after_ms: 1_000
|
||||
}
|
||||
))
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
store
|
||||
.consume_token(
|
||||
"valkey-retry-after",
|
||||
2_000_000,
|
||||
2_000_000,
|
||||
0,
|
||||
Duration::from_secs(30),
|
||||
)
|
||||
.await
|
||||
.unwrap(),
|
||||
RateLimitDecision::Allowed
|
||||
);
|
||||
assert_eq!(
|
||||
store
|
||||
.consume_token(
|
||||
"valkey-retry-after",
|
||||
2_000_000,
|
||||
2_000_000,
|
||||
0,
|
||||
Duration::from_secs(30),
|
||||
)
|
||||
.await
|
||||
.unwrap(),
|
||||
RateLimitDecision::Allowed
|
||||
);
|
||||
assert_eq!(
|
||||
store
|
||||
.consume_token(
|
||||
"valkey-retry-after",
|
||||
2_000_000,
|
||||
2_000_000,
|
||||
0,
|
||||
Duration::from_secs(30),
|
||||
)
|
||||
.await
|
||||
.unwrap(),
|
||||
RateLimitDecision::Rejected {
|
||||
retry_after_ms: 500
|
||||
}
|
||||
);
|
||||
}
|
||||
@@ -1,9 +1,14 @@
|
||||
use std::time::Duration;
|
||||
use std::{
|
||||
ffi::OsString,
|
||||
sync::{Mutex, MutexGuard},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use crank_core::{
|
||||
CacheBackend, CacheScope, CacheStoreError, CachedHeader, CachedResponse,
|
||||
CoordinationStateStore, CoordinationStateValue, RateLimitBucketState, RateLimitStateStore,
|
||||
ReplayGuardStatus, ReplayGuardStore, ResponseCacheStore,
|
||||
CoordinationStateReservation, CoordinationStateStore, CoordinationStateValue,
|
||||
RateLimitBucketState, RateLimitStateStore, ReplayGuardStatus, ReplayGuardStore,
|
||||
ResponseCacheStore,
|
||||
};
|
||||
use crank_runtime::{
|
||||
InMemoryCoordinationStateStore, InMemoryRateLimitStateStore, InMemoryReplayGuardStore,
|
||||
@@ -12,6 +17,13 @@ use crank_runtime::{
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
const CACHE_ENV_NAMES: [&str; 3] = [
|
||||
"CRANK_CACHE_BACKEND",
|
||||
"CRANK_CACHE_URL",
|
||||
"CRANK_CACHE_DEFAULT_TTL_MS",
|
||||
];
|
||||
static CACHE_ENV_LOCK: Mutex<()> = Mutex::new(());
|
||||
|
||||
#[test]
|
||||
fn defaults_to_in_memory_cache_without_url() {
|
||||
let config = RuntimeCacheConfig::default();
|
||||
@@ -21,8 +33,24 @@ fn defaults_to_in_memory_cache_without_url() {
|
||||
assert_eq!(config.default_ttl_ms, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn treats_blank_optional_cache_values_as_unset() {
|
||||
let _env = IsolatedCacheEnv::new();
|
||||
unsafe {
|
||||
std::env::set_var("CRANK_CACHE_URL", " ");
|
||||
std::env::set_var("CRANK_CACHE_DEFAULT_TTL_MS", " ");
|
||||
}
|
||||
|
||||
let config = RuntimeCacheConfig::from_env().unwrap();
|
||||
|
||||
assert_eq!(config.backend, CacheBackend::Memory);
|
||||
assert_eq!(config.url, None);
|
||||
assert_eq!(config.default_ttl_ms, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn loads_valkey_config_from_env() {
|
||||
let _env = IsolatedCacheEnv::new();
|
||||
unsafe {
|
||||
std::env::set_var("CRANK_CACHE_BACKEND", "valkey");
|
||||
std::env::set_var("CRANK_CACHE_URL", "redis://cache:6379/0");
|
||||
@@ -34,16 +62,11 @@ fn loads_valkey_config_from_env() {
|
||||
assert_eq!(config.backend, CacheBackend::Valkey);
|
||||
assert_eq!(config.url.as_deref(), Some("redis://cache:6379/0"));
|
||||
assert_eq!(config.default_ttl_ms, Some(15_000));
|
||||
|
||||
unsafe {
|
||||
std::env::remove_var("CRANK_CACHE_BACKEND");
|
||||
std::env::remove_var("CRANK_CACHE_URL");
|
||||
std::env::remove_var("CRANK_CACHE_DEFAULT_TTL_MS");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_external_backend_without_url() {
|
||||
let _env = IsolatedCacheEnv::new();
|
||||
unsafe {
|
||||
std::env::set_var("CRANK_CACHE_BACKEND", "redis");
|
||||
std::env::remove_var("CRANK_CACHE_URL");
|
||||
@@ -57,14 +80,11 @@ fn rejects_external_backend_without_url() {
|
||||
backend: CacheBackend::Redis
|
||||
}
|
||||
));
|
||||
|
||||
unsafe {
|
||||
std::env::remove_var("CRANK_CACHE_BACKEND");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_zero_ttl() {
|
||||
let _env = IsolatedCacheEnv::new();
|
||||
unsafe {
|
||||
std::env::set_var("CRANK_CACHE_DEFAULT_TTL_MS", "0");
|
||||
}
|
||||
@@ -77,9 +97,44 @@ fn rejects_zero_ttl() {
|
||||
name: "CRANK_CACHE_DEFAULT_TTL_MS"
|
||||
}
|
||||
));
|
||||
}
|
||||
|
||||
unsafe {
|
||||
std::env::remove_var("CRANK_CACHE_DEFAULT_TTL_MS");
|
||||
struct IsolatedCacheEnv {
|
||||
_lock: MutexGuard<'static, ()>,
|
||||
previous: Vec<(&'static str, Option<OsString>)>,
|
||||
}
|
||||
|
||||
impl IsolatedCacheEnv {
|
||||
fn new() -> Self {
|
||||
let lock = CACHE_ENV_LOCK
|
||||
.lock()
|
||||
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
||||
let previous = CACHE_ENV_NAMES
|
||||
.iter()
|
||||
.map(|name| (*name, std::env::var_os(name)))
|
||||
.collect();
|
||||
for name in CACHE_ENV_NAMES {
|
||||
unsafe {
|
||||
std::env::remove_var(name);
|
||||
}
|
||||
}
|
||||
Self {
|
||||
_lock: lock,
|
||||
previous,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for IsolatedCacheEnv {
|
||||
fn drop(&mut self) {
|
||||
for (name, value) in &self.previous {
|
||||
unsafe {
|
||||
match value {
|
||||
Some(value) => std::env::set_var(name, value),
|
||||
None => std::env::remove_var(name),
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -209,6 +264,81 @@ async fn in_memory_coordination_store_scopes_keys() {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn in_memory_coordination_store_atomically_takes_and_reserves_values() {
|
||||
let store = InMemoryCoordinationStateStore::default();
|
||||
let value = CoordinationStateValue {
|
||||
payload: json!({ "state": "pending" }),
|
||||
};
|
||||
store
|
||||
.put_value(
|
||||
CacheScope::Coordination,
|
||||
"atomic-job",
|
||||
value.clone(),
|
||||
Duration::from_secs(20),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let (first, second) = tokio::join!(
|
||||
store.take_value(CacheScope::Coordination, "atomic-job"),
|
||||
store.take_value(CacheScope::Coordination, "atomic-job")
|
||||
);
|
||||
assert_eq!(
|
||||
usize::from(first.unwrap().is_some()) + usize::from(second.unwrap().is_some()),
|
||||
1
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
store
|
||||
.reserve_value(
|
||||
CacheScope::Coordination,
|
||||
"reservation",
|
||||
value.clone(),
|
||||
Duration::from_secs(20),
|
||||
)
|
||||
.await
|
||||
.unwrap(),
|
||||
CoordinationStateReservation::Reserved
|
||||
);
|
||||
assert_eq!(
|
||||
store
|
||||
.reserve_value(
|
||||
CacheScope::Coordination,
|
||||
"reservation",
|
||||
CoordinationStateValue {
|
||||
payload: json!({ "state": "other" }),
|
||||
},
|
||||
Duration::from_secs(20),
|
||||
)
|
||||
.await
|
||||
.unwrap(),
|
||||
CoordinationStateReservation::Existing(value.clone())
|
||||
);
|
||||
let completed = CoordinationStateValue {
|
||||
payload: json!({ "state": "completed" }),
|
||||
};
|
||||
assert!(
|
||||
store
|
||||
.compare_and_set_value(
|
||||
CacheScope::Coordination,
|
||||
"reservation",
|
||||
&value,
|
||||
completed.clone(),
|
||||
Duration::from_secs(20),
|
||||
)
|
||||
.await
|
||||
.unwrap()
|
||||
);
|
||||
assert_eq!(
|
||||
store
|
||||
.get_value(CacheScope::Coordination, "reservation")
|
||||
.await
|
||||
.unwrap(),
|
||||
Some(completed)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn in_memory_stores_reject_empty_keys() {
|
||||
let response_store = InMemoryResponseCacheStore::default();
|
||||
|
||||
@@ -3,6 +3,7 @@ name = "crank-schema"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
rust-version.workspace = true
|
||||
publish.workspace = true
|
||||
version.workspace = true
|
||||
|
||||
[dependencies]
|
||||
|
||||
@@ -3,6 +3,7 @@ name = "crank-test-support"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
rust-version.workspace = true
|
||||
publish.workspace = true
|
||||
version.workspace = true
|
||||
|
||||
[dependencies]
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
[package]
|
||||
name = "crank-trace"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
rust-version.workspace = true
|
||||
publish.workspace = true
|
||||
version.workspace = true
|
||||
|
||||
[lib]
|
||||
path = "src/lib.rs"
|
||||
|
||||
[dependencies]
|
||||
tracing.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
tracing-subscriber.workspace = true
|
||||
@@ -0,0 +1,198 @@
|
||||
//! Закрытый семантический контракт spans Crank.
|
||||
//!
|
||||
//! Этот crate не настраивает subscriber и не знает об OTLP. Он ограничивает
|
||||
//! имена и атрибуты стадий статическим словарём, чтобы продуктовые crate не
|
||||
//! могли случайно экспортировать пользовательские данные.
|
||||
|
||||
use std::future::Future;
|
||||
|
||||
use tracing::{Instrument, Span, field::Empty, info_span};
|
||||
|
||||
macro_rules! stage_span {
|
||||
($name:literal) => {
|
||||
info_span!(
|
||||
target: "crank::trace",
|
||||
$name,
|
||||
outcome = Empty,
|
||||
error.category = Empty,
|
||||
)
|
||||
};
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub enum Stage {
|
||||
McpRateLimit,
|
||||
McpAccessCheck,
|
||||
McpCatalogLoad,
|
||||
McpToolsResolve,
|
||||
ApprovalCheck,
|
||||
RuntimeExecute,
|
||||
RuntimeArgumentsMap,
|
||||
RuntimeIdempotency,
|
||||
UpstreamHttp,
|
||||
RuntimeResponseTransform,
|
||||
AuthResolve,
|
||||
ApprovalRecovery,
|
||||
HistoryWrite,
|
||||
DbQuery,
|
||||
}
|
||||
|
||||
impl Stage {
|
||||
pub fn span(self) -> Span {
|
||||
match self {
|
||||
Self::McpRateLimit => stage_span!("mcp.rate_limit"),
|
||||
Self::McpAccessCheck => stage_span!("mcp.access.check"),
|
||||
Self::McpCatalogLoad => stage_span!("mcp.catalog.load"),
|
||||
Self::McpToolsResolve => stage_span!("mcp.tools.resolve"),
|
||||
Self::ApprovalCheck => stage_span!("approval.check"),
|
||||
Self::RuntimeExecute => stage_span!("runtime.execute"),
|
||||
Self::RuntimeArgumentsMap => stage_span!("runtime.arguments.map"),
|
||||
Self::RuntimeIdempotency => stage_span!("runtime.idempotency"),
|
||||
Self::UpstreamHttp => stage_span!("upstream.http"),
|
||||
Self::RuntimeResponseTransform => stage_span!("runtime.response.transform"),
|
||||
Self::AuthResolve => stage_span!("auth.resolve"),
|
||||
Self::ApprovalRecovery => stage_span!("approval.recovery"),
|
||||
Self::HistoryWrite => stage_span!("history.write"),
|
||||
Self::DbQuery => info_span!(
|
||||
target: "crank::trace",
|
||||
"db.query",
|
||||
outcome = Empty,
|
||||
error.category = Empty,
|
||||
db.system = "postgresql",
|
||||
db.operation = Empty,
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn db_span(self, operation: DbOperation) -> Option<Span> {
|
||||
if self != Self::DbQuery {
|
||||
return None;
|
||||
}
|
||||
let span = self.span();
|
||||
span.record("db.operation", operation.as_str());
|
||||
Some(span)
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn observe_db_query<T, E>(
|
||||
operation: DbOperation,
|
||||
future: impl Future<Output = Result<T, E>>,
|
||||
) -> Result<T, E> {
|
||||
let span = Stage::DbQuery
|
||||
.db_span(operation)
|
||||
.expect("database operation requires db.query stage");
|
||||
let result = future.instrument(span.clone()).await;
|
||||
match &result {
|
||||
Ok(_) => StageOutcome::Success.record(&span),
|
||||
Err(_) => {
|
||||
StageOutcome::Error.record(&span);
|
||||
ErrorCategory::Database.record(&span);
|
||||
}
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub enum StageOutcome {
|
||||
Success,
|
||||
Error,
|
||||
Allowed,
|
||||
Denied,
|
||||
Required,
|
||||
Replay,
|
||||
Execute,
|
||||
Skipped,
|
||||
CacheHit,
|
||||
}
|
||||
|
||||
impl StageOutcome {
|
||||
pub fn record(self, span: &Span) {
|
||||
span.record("outcome", self.as_str());
|
||||
}
|
||||
|
||||
pub const fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Success => "success",
|
||||
Self::Error => "error",
|
||||
Self::Allowed => "allowed",
|
||||
Self::Denied => "denied",
|
||||
Self::Required => "required",
|
||||
Self::Replay => "replay",
|
||||
Self::Execute => "execute",
|
||||
Self::Skipped => "skipped",
|
||||
Self::CacheHit => "cache_hit",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub enum ErrorCategory {
|
||||
Access,
|
||||
RateLimit,
|
||||
Catalog,
|
||||
Approval,
|
||||
Idempotency,
|
||||
Schema,
|
||||
Mapping,
|
||||
Upstream,
|
||||
Transformation,
|
||||
History,
|
||||
Database,
|
||||
Concurrency,
|
||||
Configuration,
|
||||
Internal,
|
||||
}
|
||||
|
||||
impl ErrorCategory {
|
||||
pub fn record(self, span: &Span) {
|
||||
span.record("error.category", self.as_str());
|
||||
}
|
||||
|
||||
pub const fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Access => "access",
|
||||
Self::RateLimit => "rate_limit",
|
||||
Self::Catalog => "catalog",
|
||||
Self::Approval => "approval",
|
||||
Self::Idempotency => "idempotency",
|
||||
Self::Schema => "schema",
|
||||
Self::Mapping => "mapping",
|
||||
Self::Upstream => "upstream",
|
||||
Self::Transformation => "transformation",
|
||||
Self::History => "history",
|
||||
Self::Database => "database",
|
||||
Self::Concurrency => "concurrency",
|
||||
Self::Configuration => "configuration",
|
||||
Self::Internal => "internal",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub enum DbOperation {
|
||||
MachineAccessRead,
|
||||
MachineAccessTouch,
|
||||
CatalogLoad,
|
||||
ApprovalRead,
|
||||
ApprovalWrite,
|
||||
AuthProfileRead,
|
||||
SecretRead,
|
||||
SecretTouch,
|
||||
InvocationHistoryWrite,
|
||||
}
|
||||
|
||||
impl DbOperation {
|
||||
pub const fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::MachineAccessRead => "machine_access.read",
|
||||
Self::MachineAccessTouch => "machine_access.touch",
|
||||
Self::CatalogLoad => "catalog.load",
|
||||
Self::ApprovalRead => "approval.read",
|
||||
Self::ApprovalWrite => "approval.write",
|
||||
Self::AuthProfileRead => "auth_profile.read",
|
||||
Self::SecretRead => "secret.read",
|
||||
Self::SecretTouch => "secret.touch",
|
||||
Self::InvocationHistoryWrite => "invocation_history.write",
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,106 @@
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use crank_trace::{DbOperation, ErrorCategory, Stage, StageOutcome};
|
||||
use tracing::{Id, Subscriber, field::Visit};
|
||||
use tracing_subscriber::{Layer, layer::SubscriberExt, registry::LookupSpan};
|
||||
|
||||
#[test]
|
||||
fn stage_names_and_attributes_are_closed() {
|
||||
let captured = Arc::new(Mutex::new(Vec::new()));
|
||||
let subscriber = tracing_subscriber::registry().with(CaptureLayer(Arc::clone(&captured)));
|
||||
|
||||
tracing::subscriber::with_default(subscriber, || {
|
||||
let span = Stage::RuntimeExecute.span();
|
||||
StageOutcome::Success.record(&span);
|
||||
ErrorCategory::Mapping.record(&span);
|
||||
drop(span);
|
||||
|
||||
let db_span = Stage::DbQuery
|
||||
.db_span(DbOperation::InvocationHistoryWrite)
|
||||
.expect("db stage accepts a db operation");
|
||||
StageOutcome::Error.record(&db_span);
|
||||
drop(db_span);
|
||||
});
|
||||
|
||||
let spans = captured.lock().expect("captured spans");
|
||||
assert_eq!(spans[0].name, "runtime.execute");
|
||||
assert_eq!(spans[0].fields["outcome"], "success");
|
||||
assert_eq!(spans[0].fields["error.category"], "mapping");
|
||||
assert_eq!(spans[1].name, "db.query");
|
||||
assert_eq!(spans[1].fields["db.system"], "postgresql");
|
||||
assert_eq!(spans[1].fields["db.operation"], "invocation_history.write");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn non_database_stage_rejects_database_attributes() {
|
||||
assert!(
|
||||
Stage::RuntimeExecute
|
||||
.db_span(DbOperation::CatalogLoad)
|
||||
.is_none()
|
||||
);
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct CaptureLayer(Arc<Mutex<Vec<CapturedSpan>>>);
|
||||
|
||||
#[derive(Debug)]
|
||||
struct CapturedSpan {
|
||||
name: &'static str,
|
||||
fields: std::collections::BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
impl<S> Layer<S> for CaptureLayer
|
||||
where
|
||||
S: Subscriber + for<'lookup> LookupSpan<'lookup>,
|
||||
{
|
||||
fn on_new_span(
|
||||
&self,
|
||||
attributes: &tracing::span::Attributes<'_>,
|
||||
id: &Id,
|
||||
context: tracing_subscriber::layer::Context<'_, S>,
|
||||
) {
|
||||
let mut visitor = FieldVisitor::default();
|
||||
attributes.record(&mut visitor);
|
||||
context
|
||||
.span(id)
|
||||
.expect("span exists")
|
||||
.extensions_mut()
|
||||
.insert(self.0.lock().expect("capture lock").len());
|
||||
self.0.lock().expect("capture lock").push(CapturedSpan {
|
||||
name: attributes.metadata().name(),
|
||||
fields: visitor.fields,
|
||||
});
|
||||
}
|
||||
|
||||
fn on_record(
|
||||
&self,
|
||||
id: &Id,
|
||||
values: &tracing::span::Record<'_>,
|
||||
context: tracing_subscriber::layer::Context<'_, S>,
|
||||
) {
|
||||
let span = context.span(id).expect("span exists");
|
||||
let index = *span.extensions().get::<usize>().expect("capture index");
|
||||
let mut visitor = FieldVisitor::default();
|
||||
values.record(&mut visitor);
|
||||
self.0.lock().expect("capture lock")[index]
|
||||
.fields
|
||||
.extend(visitor.fields);
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct FieldVisitor {
|
||||
fields: std::collections::BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
impl Visit for FieldVisitor {
|
||||
fn record_str(&mut self, field: &tracing::field::Field, value: &str) {
|
||||
self.fields
|
||||
.insert(field.name().to_owned(), value.to_owned());
|
||||
}
|
||||
|
||||
fn record_debug(&mut self, field: &tracing::field::Field, value: &dyn std::fmt::Debug) {
|
||||
self.fields
|
||||
.insert(field.name().to_owned(), format!("{value:?}"));
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user