diff --git a/.github/workflows/go.yml b/.github/workflows/go.yml index 6fdd3ef..b7004b1 100644 --- a/.github/workflows/go.yml +++ b/.github/workflows/go.yml @@ -16,7 +16,7 @@ jobs: - name: Set up Go uses: actions/setup-go@v5 with: - go-version: '1.22' + go-version: '1.25' - name: Cache Go modules uses: actions/cache@v4 @@ -41,7 +41,7 @@ jobs: - name: Set up Go uses: actions/setup-go@v5 with: - go-version: '1.22' + go-version: '1.25' - name: Cache Go modules uses: actions/cache@v4 @@ -66,17 +66,17 @@ jobs: - name: Set up Go uses: actions/setup-go@v5 with: - go-version: '1.22' + go-version: '1.25' - name: golangci-lint (clients/go) - uses: golangci/golangci-lint-action@v6 + uses: golangci/golangci-lint-action@v9 with: version: latest working-directory: clients/go args: --timeout=5m - name: golangci-lint (gateway) - uses: golangci/golangci-lint-action@v6 + uses: golangci/golangci-lint-action@v9 with: version: latest working-directory: gateway diff --git a/.gitignore b/.gitignore index b8581b0..ff8ff27 100644 --- a/.gitignore +++ b/.gitignore @@ -36,6 +36,8 @@ pnpm-debug.log* dist/ build/ .next/ +.next-build/ +.next-playwright-*/ out/ .turbo/ .vercel/ diff --git a/CHANGELOG.md b/CHANGELOG.md new file mode 100644 index 0000000..24a3b5a --- /dev/null +++ b/CHANGELOG.md @@ -0,0 +1,30 @@ +# Changelog + +All notable changes to this project are documented in this file. + +## [Unreleased] + +### Added + +- Add authenticated Streamable HTTP MCP with scoped read and write tools. +- Add OAuth 2.1 with PKCE, dynamic client registration, and browser OIDC. +- Add personal API tokens with scopes, expiry, one-time display, and revocation. +- Add OpenAPI JSON and YAML documents, `llms.txt`, and client authentication guidance. +- Add batch append and blob reads, bounded trace pages, exact turn hydration, and faster typed projection. +- Add responsive token management and debugger views for mobile devices. + +### Changed + +- Enforce write scopes and stricter proxy-header handling. +- Correct separate tool-result hydration and named-key MessagePack projection. Numeric field tags continue to have priority. + +### Fixed + +- Pin pnpm 9 in Node 20 container builds. + +### Compatibility + +- Gateway deployments must now provide a `SESSION_SECRET` of at least 32 bytes. +- Non-GET API requests now require authentication with the `cxdb:write` scope. +- Scoped credentials used with context, metrics, and event reads must include `cxdb:read`. Existing browser sessions and built-in service credentials receive both scopes. +- There are no known breaking stored-data changes. diff --git a/Dockerfile b/Dockerfile index caf2e08..0df63a7 100644 --- a/Dockerfile +++ b/Dockerfile @@ -11,7 +11,7 @@ FROM node:20-alpine AS frontend WORKDIR /app # Install pnpm -RUN corepack enable && corepack prepare pnpm@latest --activate +RUN corepack enable && corepack prepare pnpm@9.15.9 --activate # Copy package files COPY frontend/package.json frontend/pnpm-lock.yaml* ./ @@ -28,7 +28,7 @@ RUN pnpm build # ============================================ # Stage 2: Build Rust binary # ============================================ -FROM rust:1.92-bookworm AS backend +FROM rust:1.94.1-bookworm AS backend WORKDIR /app diff --git a/clients/go/cmd/cxdb-blob-verify/main.go b/clients/go/cmd/cxdb-blob-verify/main.go index 80e1be1..4a56c87 100644 --- a/clients/go/cmd/cxdb-blob-verify/main.go +++ b/clients/go/cmd/cxdb-blob-verify/main.go @@ -69,7 +69,7 @@ func main() { // --- Test 1: PutBlob + GetBlob round-trip --- run("PutBlob + GetBlob round-trip", func() error { client := dial() - defer client.Close() + defer func() { _ = client.Close() }() ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) defer cancel() @@ -100,7 +100,7 @@ func main() { // --- Test 2: GetBlob nonexistent hash returns ErrBlobNotFound --- run("GetBlob nonexistent hash (expect ErrBlobNotFound)", func() error { client := dial() - defer client.Close() + defer func() { _ = client.Close() }() ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) defer cancel() @@ -123,7 +123,7 @@ func main() { // --- Test 3: PutBlobIfAbsent deduplication --- run("PutBlobIfAbsent deduplication", func() error { client := dial() - defer client.Close() + defer func() { _ = client.Close() }() ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) defer cancel() @@ -166,7 +166,7 @@ func main() { // --- Test 4: Large blob (1 MiB) --- run("Large blob (1 MiB) round-trip", func() error { client := dial() - defer client.Close() + defer func() { _ = client.Close() }() ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() @@ -196,7 +196,7 @@ func main() { // --- Test 5: Connection survives a not-found --- run("Connection survives not-found", func() error { client := dial() - defer client.Close() + defer func() { _ = client.Close() }() ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) defer cancel() @@ -244,7 +244,7 @@ func main() { if err != nil { return fmt.Errorf("DialReconnecting: %w", err) } - defer rc.Close() + defer func() { _ = rc.Close() }() ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) defer cancel() diff --git a/clients/go/fs_test.go b/clients/go/fs_test.go index b354c8e..0770b73 100644 --- a/clients/go/fs_test.go +++ b/clients/go/fs_test.go @@ -46,7 +46,7 @@ func mockServer(t *testing.T, handler mockHandler) (addr string, cleanup func()) if err != nil { return // listener closed } - defer conn.Close() + defer func() { _ = conn.Close() }() // --- HELLO handshake --- helloFrame, err := mockReadFrame(conn) @@ -107,7 +107,7 @@ func mockServer(t *testing.T, handler mockHandler) (addr string, cleanup func()) }() cleanup = func() { - ln.Close() + _ = ln.Close() wg.Wait() } t.Cleanup(cleanup) @@ -180,7 +180,7 @@ func TestGetBlob_Success(t *testing.T) { if err != nil { t.Fatalf("Dial: %v", err) } - defer client.Close() + defer func() { _ = client.Close() }() data, err := client.GetBlob(context.Background(), requestHash) if err != nil { @@ -200,7 +200,7 @@ func TestGetBlob_NotFound(t *testing.T) { if err != nil { t.Fatalf("Dial: %v", err) } - defer client.Close() + defer func() { _ = client.Close() }() var hash [32]byte _, err = client.GetBlob(context.Background(), hash) @@ -222,7 +222,7 @@ func TestGetBlob_ResponseTooShort(t *testing.T) { if err != nil { t.Fatalf("Dial: %v", err) } - defer client.Close() + defer func() { _ = client.Close() }() var hash [32]byte _, err = client.GetBlob(context.Background(), hash) @@ -247,7 +247,7 @@ func TestGetBlob_PayloadTruncated(t *testing.T) { if err != nil { t.Fatalf("Dial: %v", err) } - defer client.Close() + defer func() { _ = client.Close() }() var hash [32]byte _, err = client.GetBlob(context.Background(), hash) @@ -323,7 +323,7 @@ func TestPutBlobThenGetBlob_Roundtrip(t *testing.T) { if err != nil { t.Fatalf("Dial: %v", err) } - defer client.Close() + defer func() { _ = client.Close() }() blobData := []byte("the quick brown fox jumps over the lazy dog") diff --git a/clients/go/fstree/integration_test.go b/clients/go/fstree/integration_test.go index 244db28..b817ff7 100644 --- a/clients/go/fstree/integration_test.go +++ b/clients/go/fstree/integration_test.go @@ -44,7 +44,7 @@ func TestE2E_FilesystemSnapshots(t *testing.T) { if err != nil { t.Fatalf("Failed to connect to server at %s: %v\nMake sure the server is running", binaryAddr, err) } - defer client.Close() + defer func() { _ = client.Close() }() t.Logf("Connected to CXDB server, session ID: %d", client.SessionID()) @@ -483,9 +483,9 @@ func makePayload(t *testing.T, itemType, text string) []byte { t.Helper() item := map[uint64]any{ - 1: itemType, // type - 2: "complete", // status - 3: time.Now().UnixMilli(), // timestamp + 1: itemType, // type + 2: "complete", // status + 3: time.Now().UnixMilli(), // timestamp 4: fmt.Sprintf("test-%d", time.Now().UnixNano()), // id } @@ -520,7 +520,7 @@ func verifyHTTPFsListing(t *testing.T, turnID uint64, path string, expectedNames if err != nil { t.Fatalf("HTTP GET %s failed: %v", url, err) } - defer resp.Body.Close() + defer func() { _ = resp.Body.Close() }() if resp.StatusCode != 200 { body, _ := io.ReadAll(resp.Body) @@ -558,7 +558,7 @@ func verifyHTTPFsFileContent(t *testing.T, turnID uint64, path string, expectedC if err != nil { t.Fatalf("HTTP GET %s failed: %v", url, err) } - defer resp.Body.Close() + defer func() { _ = resp.Body.Close() }() if resp.StatusCode != 200 { body, _ := io.ReadAll(resp.Body) @@ -591,7 +591,7 @@ func verifyHTTPFsFileNotFound(t *testing.T, turnID uint64, path string) { if err != nil { t.Fatalf("HTTP GET %s failed: %v", url, err) } - defer resp.Body.Close() + defer func() { _ = resp.Body.Close() }() if resp.StatusCode != 404 { t.Errorf("Expected 404 for turn %d path '%s', got %d", turnID, path, resp.StatusCode) @@ -615,7 +615,7 @@ func TestE2E_BlobDeduplication(t *testing.T) { if err != nil { t.Fatalf("Failed to connect: %v", err) } - defer client.Close() + defer func() { _ = client.Close() }() // Create unique content uniqueContent := []byte(fmt.Sprintf("unique content %d", time.Now().UnixNano())) @@ -655,7 +655,7 @@ func TestE2E_FsRootInheritance(t *testing.T) { if err != nil { t.Fatalf("Failed to connect: %v", err) } - defer client.Close() + defer func() { _ = client.Close() }() workDir := t.TempDir() os.WriteFile(filepath.Join(workDir, "test.txt"), []byte("inherited content"), 0644) diff --git a/clients/rust/src/follow.rs b/clients/rust/src/follow.rs index 4a28f86..6685f24 100644 --- a/clients/rust/src/follow.rs +++ b/clients/rust/src/follow.rs @@ -452,6 +452,16 @@ mod tests { follow_turns(&ctx, event_rx, client.clone(), vec![with_follow_buffer(10)]); event_tx.send(make_turn_event(context_id, 2, 1)).unwrap(); + let mut got = vec![ + out.recv_timeout(Duration::from_secs(1)) + .expect("backfill turn 1") + .turn + .turn_id, + out.recv_timeout(Duration::from_secs(1)) + .expect("backfill turn 2") + .turn + .turn_id, + ]; client.set_context( context_id, @@ -495,7 +505,7 @@ mod tests { event_tx.send(make_turn_event(context_id, 3, 2)).unwrap(); drop(event_tx); - let got: Vec = out.iter().map(|turn| turn.turn.turn_id).collect(); + got.extend(out.iter().map(|turn| turn.turn.turn_id)); if let Some(err) = errs.try_iter().next() { panic!("unexpected error: {}", err); } diff --git a/cxtx/src/delivery.rs b/cxtx/src/delivery.rs index 3803a8f..3418ea5 100644 --- a/cxtx/src/delivery.rs +++ b/cxtx/src/delivery.rs @@ -1,3 +1,6 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + use anyhow::{anyhow, Result}; use std::collections::VecDeque; use std::time::Duration; @@ -28,7 +31,7 @@ enum WorkerMessage { #[derive(Debug, Clone)] enum QueueItem { CreateContext, - Append(TurnEnvelope), + Append(Box), } impl DeliveryHandle { @@ -54,7 +57,7 @@ impl DeliveryHandle { pub async fn enqueue_turn(&self, turn: TurnEnvelope) -> Result<()> { self.tx - .send(WorkerMessage::Enqueue(QueueItem::Append(turn))) + .send(WorkerMessage::Enqueue(QueueItem::Append(Box::new(turn)))) .await .map_err(|_| anyhow!("delivery worker is no longer running")) } @@ -135,8 +138,9 @@ impl DeliveryWorker { if self.degraded && self.queue.is_empty() && !self.recovery_turn_enqueued { self.recovery_turn_enqueued = true; - self.queue - .push_back(QueueItem::Append(self.session.ingest_recovered_turn(0))); + self.queue.push_back(QueueItem::Append(Box::new( + self.session.ingest_recovered_turn(0), + ))); } else if self.degraded && self.recovery_turn_enqueued && matches!(item, QueueItem::Append(_)) @@ -223,9 +227,9 @@ impl DeliveryWorker { self.degraded = true; self.recovery_turn_enqueued = false; - self.queue.push_back(QueueItem::Append( + self.queue.push_back(QueueItem::Append(Box::new( self.session.ingest_degraded_turn(self.queue.len(), error), - )); + ))); eprintln!("cxtx: CXDB ingest unavailable, entering queued-delivery mode"); } diff --git a/cxtx/src/provider/anthropic.rs b/cxtx/src/provider/anthropic.rs index f36824c..f5686ef 100644 --- a/cxtx/src/provider/anthropic.rs +++ b/cxtx/src/provider/anthropic.rs @@ -1,3 +1,6 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + use serde_json::Value; use std::collections::BTreeMap; @@ -425,7 +428,8 @@ fn is_bootstrap_user_block(block: &Value) -> bool { text.starts_with("") && (text.contains("SessionStart hook additional context") || text.contains("The following skills are available for use with the Skill tool:") - || text.contains("As you answer the user's questions, you can use the following context:")) + || text + .contains("As you answer the user's questions, you can use the following context:")) } fn parse_assistant_content( diff --git a/cxtx/src/provider/mod.rs b/cxtx/src/provider/mod.rs index f573615..0540077 100644 --- a/cxtx/src/provider/mod.rs +++ b/cxtx/src/provider/mod.rs @@ -1,3 +1,6 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + pub mod anthropic; pub mod openai; @@ -69,10 +72,7 @@ impl ProviderKind { out.extend(["-c".to_string(), "prefer_websockets=false".to_string()]); } if !has_codex_feature_override(args, "responses_websockets") { - out.extend([ - "--disable".to_string(), - "responses_websockets".to_string(), - ]); + out.extend(["--disable".to_string(), "responses_websockets".to_string()]); } if !has_codex_feature_override(args, "responses_websockets_v2") { out.extend([ diff --git a/cxtx/src/provider/openai.rs b/cxtx/src/provider/openai.rs index 142bf0c..6777330 100644 --- a/cxtx/src/provider/openai.rs +++ b/cxtx/src/provider/openai.rs @@ -1,3 +1,6 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + use serde_json::Value; use crate::provider::{ExchangeState, PreparedExchange}; @@ -418,7 +421,10 @@ fn is_bootstrap_input_item(item: &Value) -> bool { || trimmed.starts_with("") } -fn parse_assistant_payload(payload: &Value, fallback_model: Option<&str>) -> Result, String> { +fn parse_assistant_payload( + payload: &Value, + fallback_model: Option<&str>, +) -> Result, String> { if let Some(message) = payload .get("choices") .and_then(Value::as_array) @@ -561,7 +567,11 @@ fn absorb_tool_call_delta(slots: &mut Vec, delta: &Value) { } } -fn absorb_responses_output(response: &Value, content: &mut String, tool_calls: &mut Vec) { +fn absorb_responses_output( + response: &Value, + content: &mut String, + tool_calls: &mut Vec, +) { let Some(items) = response.get("output").and_then(Value::as_array) else { return; }; @@ -578,7 +588,9 @@ fn absorb_responses_output_item( match item.get("type").and_then(Value::as_str) { Some("message") => { if content.is_empty() { - content.push_str(&content_to_text(item.get("content").unwrap_or(&Value::Null))); + content.push_str(&content_to_text( + item.get("content").unwrap_or(&Value::Null), + )); } } Some("function_call") => { diff --git a/cxtx/src/proxy.rs b/cxtx/src/proxy.rs index fa429db..7836ab7 100644 --- a/cxtx/src/proxy.rs +++ b/cxtx/src/proxy.rs @@ -1,3 +1,6 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + use anyhow::{anyhow, Context, Result}; use async_stream::stream; use axum::body::{to_bytes, Body}; @@ -13,11 +16,11 @@ use serde_json::Value; use std::net::SocketAddr; use std::sync::Arc; use tokio::net::TcpListener; +use tokio::sync::{mpsc, oneshot, RwLock}; use tokio_tungstenite::tungstenite::client::IntoClientRequest; use tokio_tungstenite::tungstenite::error::ProtocolError as TungsteniteProtocolError; use tokio_tungstenite::tungstenite::protocol::Message as UpstreamWsMessage; use tokio_tungstenite::tungstenite::Error as TungsteniteError; -use tokio::sync::{mpsc, oneshot, RwLock}; use url::Url; use crate::delivery::DeliveryHandle; @@ -317,7 +320,9 @@ async fn handle_websocket_proxy_request( let status = StatusCode::from_u16(upstream_response.status().as_u16()) .unwrap_or(StatusCode::SWITCHING_PROTOCOLS); - let request_id = state.provider.request_id_from_headers(upstream_response.headers()); + let request_id = state + .provider + .request_id_from_headers(upstream_response.headers()); let response_artifact = state .ledger .record_response( @@ -363,6 +368,7 @@ async fn handle_websocket_proxy_request( .into_response()) } +#[allow(clippy::too_many_arguments)] async fn body_response( state: ProxyState, delivery: DeliveryHandle, @@ -412,6 +418,7 @@ async fn body_response( .context("failed to build proxied response") } +#[allow(clippy::too_many_arguments)] async fn stream_response( state: ProxyState, delivery: DeliveryHandle, @@ -452,13 +459,16 @@ async fn stream_response( stream_artifact_refs.stream_path = Some(path); } Err(err) => { - delivery.enqueue_turn(session.provider_error_turn( - &exchange_id, - "artifact_write_error", - &format!("failed to persist stream frame: {err}"), - request_id_for_stream.as_deref(), - &stream_artifact_refs, - )).await.ok(); + delivery + .enqueue_turn(session.provider_error_turn( + &exchange_id, + "artifact_write_error", + &format!("failed to persist stream frame: {err}"), + request_id_for_stream.as_deref(), + &stream_artifact_refs, + )) + .await + .ok(); } } exchange_state.absorb_sse_frame(&frame); @@ -468,40 +478,47 @@ async fn stream_response( } } Err(err) => { - delivery.enqueue_turn(session.provider_error_turn( - &exchange_id, - "stream_transport_error", - &format!("failed to read upstream stream: {err}"), - request_id_for_stream.as_deref(), - &stream_artifact_refs, - )).await.ok(); - let _ = tx - .send(Err(std::io::Error::other(err.to_string()))) - .await; + delivery + .enqueue_turn(session.provider_error_turn( + &exchange_id, + "stream_transport_error", + &format!("failed to read upstream stream: {err}"), + request_id_for_stream.as_deref(), + &stream_artifact_refs, + )) + .await + .ok(); + let _ = tx.send(Err(std::io::Error::other(err.to_string()))).await; break; } } } - match ledger.record_response( - &exchange_id, - status.as_u16(), - request_id_for_stream.as_deref(), - content_type.as_deref(), - raw_stream.as_bytes(), - None, - ).await { + match ledger + .record_response( + &exchange_id, + status.as_u16(), + request_id_for_stream.as_deref(), + content_type.as_deref(), + raw_stream.as_bytes(), + None, + ) + .await + { Ok(path) => { stream_artifact_refs.response_path = Some(path); } Err(err) => { - delivery.enqueue_turn(session.provider_error_turn( - &exchange_id, - "artifact_write_error", - &format!("failed to persist streamed response transcript: {err}"), - request_id_for_stream.as_deref(), - &stream_artifact_refs, - )).await.ok(); + delivery + .enqueue_turn(session.provider_error_turn( + &exchange_id, + "artifact_write_error", + &format!("failed to persist streamed response transcript: {err}"), + request_id_for_stream.as_deref(), + &stream_artifact_refs, + )) + .await + .ok(); } } @@ -541,6 +558,7 @@ async fn enqueue_turns(delivery: &DeliveryHandle, turns: Vec Option Option { match message { - UpstreamWsMessage::Text(text) => Some(DownstreamWsMessage::Text(text.to_string().into())), - UpstreamWsMessage::Binary(bytes) => { - Some(DownstreamWsMessage::Binary(bytes.to_vec().into())) - } - UpstreamWsMessage::Ping(bytes) => Some(DownstreamWsMessage::Ping(bytes.to_vec().into())), - UpstreamWsMessage::Pong(bytes) => Some(DownstreamWsMessage::Pong(bytes.to_vec().into())), + UpstreamWsMessage::Text(text) => Some(DownstreamWsMessage::Text(text.to_string())), + UpstreamWsMessage::Binary(bytes) => Some(DownstreamWsMessage::Binary(bytes.to_vec())), + UpstreamWsMessage::Ping(bytes) => Some(DownstreamWsMessage::Ping(bytes.to_vec())), + UpstreamWsMessage::Pong(bytes) => Some(DownstreamWsMessage::Pong(bytes.to_vec())), UpstreamWsMessage::Close(_) => Some(DownstreamWsMessage::Close(None)), UpstreamWsMessage::Frame(_) => None, } @@ -1052,21 +1068,15 @@ mod tests { headers.insert("accept-encoding", HeaderValue::from_static("gzip, br")); let forwarded = forwardable_headers(&headers); - assert!( - !forwarded - .iter() - .any(|(name, _)| name.as_str().eq_ignore_ascii_case("host")) - ); - assert!( - !forwarded - .iter() - .any(|(name, _)| name.as_str().eq_ignore_ascii_case("accept-encoding")) - ); - assert!( - forwarded - .iter() - .any(|(name, value)| name == "authorization" && value == "Bearer test") - ); + assert!(!forwarded + .iter() + .any(|(name, _)| name.as_str().eq_ignore_ascii_case("host"))); + assert!(!forwarded + .iter() + .any(|(name, _)| name.as_str().eq_ignore_ascii_case("accept-encoding"))); + assert!(forwarded + .iter() + .any(|(name, value)| name == "authorization" && value == "Bearer test")); } #[test] @@ -1077,16 +1087,12 @@ mod tests { headers.insert("sec-websocket-version", HeaderValue::from_static("13")); let forwarded = websocket_forwardable_headers(&headers); - assert!( - forwarded - .iter() - .any(|(name, value)| name == "authorization" && value == "Bearer test") - ); - assert!( - !forwarded - .iter() - .any(|(name, _)| name.as_str().eq_ignore_ascii_case("sec-websocket-key")) - ); + assert!(forwarded + .iter() + .any(|(name, value)| name == "authorization" && value == "Bearer test")); + assert!(!forwarded + .iter() + .any(|(name, _)| name.as_str().eq_ignore_ascii_case("sec-websocket-key"))); } #[test] @@ -1111,6 +1117,7 @@ mod tests { } #[tokio::test(flavor = "multi_thread")] + #[allow(clippy::result_large_err)] async fn websocket_upgrade_requests_are_relayed_to_upstream() { let upstream_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); let upstream_addr = upstream_listener.local_addr().unwrap(); diff --git a/cxtx/tests/integration.rs b/cxtx/tests/integration.rs index aa56603..bfa5b96 100644 --- a/cxtx/tests/integration.rs +++ b/cxtx/tests/integration.rs @@ -1,3 +1,6 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + use std::collections::BTreeMap; use std::fs; use std::net::TcpListener; @@ -173,9 +176,11 @@ printf 'codex-child-stderr\n' >&2 assert!(fs::read_to_string(fixture_dir.join("openai_base_url.txt")) .unwrap() .contains("/v1")); - assert!(fs::read_to_string(fixture_dir.join("openai_base_url_env.txt")) - .unwrap() - .contains("/v1")); + assert!( + fs::read_to_string(fixture_dir.join("openai_base_url_env.txt")) + .unwrap() + .contains("/v1") + ); assert_eq!( fs::read_to_string(fixture_dir.join("openai_api_base_env.txt")).unwrap(), fs::read_to_string(fixture_dir.join("openai_api_base.txt")).unwrap() @@ -506,7 +511,13 @@ PY let item_types = turn_item_types(&turns); assert_eq!( item_types, - vec!["system", "assistant_turn", "user_input", "assistant_turn", "system"] + vec![ + "system", + "assistant_turn", + "user_input", + "assistant_turn", + "system" + ] ); assert_eq!(turns[1]["data"]["turn"]["text"], "previous answer"); assert_eq!(turns[2]["data"]["user_input"]["text"], "real prompt"); @@ -547,7 +558,10 @@ async fn websocket_proxy_uploads_canonical_turns_into_cxdb() { .unwrap(); proxy.set_delivery(delivery.clone()).await; delivery.enqueue_create_context().await.unwrap(); - delivery.enqueue_turn(session.session_start_turn()).await.unwrap(); + delivery + .enqueue_turn(session.session_start_turn()) + .await + .unwrap(); let mut proxy_url = proxy.proxy_base_url(); proxy_url.set_scheme("ws").unwrap(); @@ -1517,6 +1531,7 @@ struct MockOpenAiWebsocket { } impl MockOpenAiWebsocket { + #[allow(clippy::result_large_err)] async fn start() -> anyhow::Result { let listener = TokioTcpListener::bind("127.0.0.1:0").await?; let addr = listener.local_addr()?; @@ -2153,17 +2168,15 @@ fn turn_item_types(turns: &[Value]) -> Vec<&str> { } fn openai_input_last_user_text(payload: &Value) -> Option<&str> { - payload["input"] - .as_array() - .and_then(|items| { - items.iter().rev().find_map(|item| { - (item["role"] == "user") - .then_some(item["content"].as_array()) - .flatten() - .and_then(|content| content.last()) - .and_then(|part| part["text"].as_str()) - }) + payload["input"].as_array().and_then(|items| { + items.iter().rev().find_map(|item| { + (item["role"] == "user") + .then_some(item["content"].as_array()) + .flatten() + .and_then(|content| content.last()) + .and_then(|part| part["text"].as_str()) }) + }) } fn first_context_id(contexts: &Value) -> u64 { @@ -2193,16 +2206,14 @@ fn anthropic_last_user_text(payload: &Value) -> Option<&str> { }) } -fn wait_for_http(url: &str) -> impl std::future::Future> + '_ { - async move { - for _ in 0..40 { - match Client::new().get(url).send().await { - Ok(response) if response.status().is_success() => return Ok(()), - _ => tokio::time::sleep(Duration::from_millis(50)).await, - } +async fn wait_for_http(url: &str) -> anyhow::Result<()> { + for _ in 0..40 { + match Client::new().get(url).send().await { + Ok(response) if response.status().is_success() => return Ok(()), + _ => tokio::time::sleep(Duration::from_millis(50)).await, } - anyhow::bail!("server at {url} did not become ready") } + anyhow::bail!("server at {url} did not become ready") } fn write_executable(path: &Path, contents: &str) -> anyhow::Result<()> { diff --git a/deploy/docker-compose/.env.example b/deploy/docker-compose/.env.example index 76d3727..05b7e1d 100644 --- a/deploy/docker-compose/.env.example +++ b/deploy/docker-compose/.env.example @@ -29,6 +29,9 @@ SESSION_SECRET=your-64-character-hex-secret-here # For production: https://your-domain.com PUBLIC_BASE_URL=http://localhost:8080 +# Optional trusted reverse-proxy networks for client-specific auth rate limits. +# TRUSTED_PROXY_CIDRS=172.20.0.0/24 + # ------------------------------------------- # Gateway Configuration (Optional) # ------------------------------------------- diff --git a/docs/CQL_REFERENCE.md b/docs/CQL_REFERENCE.md new file mode 100644 index 0000000..23c948e --- /dev/null +++ b/docs/CQL_REFERENCE.md @@ -0,0 +1,87 @@ +# CQL reference + +CQL is the CXDB query language for filtering contexts through +`GET /v1/contexts/search?q={query}&limit={n}`. + +## Examples + +```text +tag = "amplifier" +tag = "amplifier" AND user = "alice" +user = "alice" +service ^= "worker" +created > "-24h" +(service = "worker" OR service = "api") AND NOT tag = "test" +``` + +## Boolean operators + +| Operator | Precedence | Example | +| --- | --- | --- | +| `NOT` | Highest | `NOT tag = "test"` | +| `AND` | Medium | `tag = "a" AND user = "b"` | +| `OR` | Lowest | `tag = "a" OR tag = "b"` | + +Use parentheses to change the normal precedence. + +## Comparison operators + +| Operator | Meaning | Example | +| --- | --- | --- | +| `=` | Exact match | `tag = "amplifier"` | +| `!=` | Not equal | `service != "test"` | +| `^=` | Starts with | `tag ^= "amp"` | +| `~=` | Case-insensitive equality | `user ~= "Alice"` | +| `^~=` | Case-insensitive prefix | `service ^~= "API"` | +| `>` | Greater than | `created > "-24h"` | +| `>=` | Greater than or equal | `depth >= 5` | +| `<` | Less than | `created < "2026-01-01"` | +| `<=` | Less than or equal | `depth <= 10` | +| `IN` | List membership | `tag IN ("a", "b")` | + +## Fields + +| Field | Type | Meaning | +| --- | --- | --- | +| `id` | number | Context ID | +| `tag` | string | Client tag | +| `title` | string | Context title | +| `label` | string | Context label | +| `user` | string | User identity | +| `service` | string | Service name | +| `host` | string | Host name | +| `trace_id` | string | Trace ID | +| `parent` | number | Parent context ID | +| `root` | number | Root context ID | +| `created` | datetime | Creation time | +| `depth` | number | Conversation depth | +| `is_live` | boolean | Active session state | + +Relative time values use `-Nh`, `-Nd`, or `-Nm`. Absolute values can use an +ISO 8601 timestamp or a date in `YYYY-MM-DD` form. + +## HTTP example + +```bash +curl --get 'http://localhost:9010/v1/contexts/search' \ + --data-urlencode 'q=tag = "amplifier" AND created > "-24h"' \ + --data-urlencode 'limit=20' +``` + +The gateway requires a signed browser session or a personal API token. + +## Grammar + +```ebnf +query = expression ; +expression = or_expr ; +or_expr = and_expr { "OR" and_expr } ; +and_expr = unary_expr { "AND" unary_expr } ; +unary_expr = [ "NOT" ] primary ; +primary = comparison | "(" expression ")" ; +comparison = field operator value ; +field = identifier ; +operator = "=" | "!=" | "^=" | "~=" | "^~=" | ">" | ">=" | "<" | "<=" | "IN" ; +value = string | number | date | list ; +list = "(" value { "," value } ")" ; +``` diff --git a/docs/client-authentication.md b/docs/client-authentication.md new file mode 100644 index 0000000..d936c98 --- /dev/null +++ b/docs/client-authentication.md @@ -0,0 +1,41 @@ +# Client authentication + +CXDB supports browser login and scoped bearer tokens through the gateway. + +## Browser login + +The gateway uses the configured OpenID Connect (OIDC) provider. The provider +must return a verified identity. The gateway creates a browser session and a +session-bound CSRF token for state-changing requests. + +## Personal API tokens + +Authenticated users can create personal tokens in the Web UI. A token has a +name, an optional expiry, and one or both of these scopes: + +- `cxdb:read` for context and turn reads. +- `cxdb:write` for context creation and turn appends. + +The token secret is shown only once. Store it in a secret manager. Send it in +the HTTP header below. Do not put a token in a URL or browser storage. + +```http +Authorization: Bearer +``` + +The Web UI uses `X-CSRF-Token` for create and revoke requests. Token metadata +does not include token secrets. A revoked or expired token cannot be used. + +## MCP OAuth + +Remote MCP clients can use OAuth 2.1 authorization code flow with PKCE. Use +the protected-resource metadata at +`/.well-known/oauth-protected-resource/mcp` and the authorization-server +metadata at `/.well-known/oauth-authorization-server`. + +The gateway delegates browser identity checks to the configured OIDC provider. +Dynamic client registration is available at `/oauth/register`. Redirect URIs +must use HTTPS or loopback HTTP. MCP clients can also use a personal bearer +token with `cxdb:read`; write tools also require `cxdb:write`. + +See [MCP guidance](mcp.md) and the published [OpenAPI JSON](../frontend/public/openapi.json). diff --git a/docs/deployment.md b/docs/deployment.md index dea8c54..ed4592f 100644 --- a/docs/deployment.md +++ b/docs/deployment.md @@ -68,6 +68,7 @@ Internet | `SESSION_SECRET` | Yes | 64-char hex string for cookie signing | | `DATABASE_PATH` | No | Session DB path (default: ./data/sessions.db) | | `ALLOWED_RENDERER_ORIGINS` | No | CSP script-src origins (comma-separated) | +| `TRUSTED_PROXY_CIDRS` | No | CIDRs of reverse proxies whose X-Forwarded-For chain can identify clients for auth rate limits | | `DEV_MODE` | No | Disable OAuth (development only) | ### Generating Secrets diff --git a/docs/http-api.md b/docs/http-api.md index c87f296..360c1c1 100644 --- a/docs/http-api.md +++ b/docs/http-api.md @@ -8,10 +8,18 @@ The CXDB HTTP gateway provides a JSON API for reading turns, managing contexts, **Development:** No authentication required when connecting directly to the Rust server -**Production:** The Go gateway provides Google OAuth authentication: -- Unauthenticated requests to `/v1/*` return `302 Found` redirect to `/login` -- After OAuth, requests include session cookie -- Session expires after 24 hours of inactivity +**Gateway authentication:** The Go gateway uses the configured OIDC provider +for browser sessions. Clients may also send a personal bearer token: + +```http +Authorization: Bearer +``` + +Read requests require `cxdb:read`. Context creation and turn append requests +also require `cxdb:write`. The Web UI creates and revokes tokens. A token +secret is shown only once. Do not put tokens in URLs or browser storage. + +See [client authentication](client-authentication.md) and [MCP guidance](mcp.md). ## Contexts diff --git a/docs/mcp.md b/docs/mcp.md new file mode 100644 index 0000000..dc8bafb --- /dev/null +++ b/docs/mcp.md @@ -0,0 +1,28 @@ +# CXDB remote MCP + +CXDB serves Streamable HTTP MCP at `/mcp`. It uses the current 2026-07-28 protocol through the official Go MCP SDK. The endpoint is stateless and validates cross-origin requests. + +## Authentication + +The gateway publishes OAuth protected-resource metadata at `/.well-known/oauth-protected-resource/mcp` and authorization-server metadata at `/.well-known/oauth-authorization-server`. + +Remote clients use OAuth 2.1 authorization code flow with PKCE S256. CXDB is the OAuth authorization server. It delegates the browser identity check to the configured OIDC provider. Dynamic client registration is available at `/oauth/register`. Redirect URIs must use HTTPS or loopback HTTP. + +Personal API tokens also work as MCP bearer tokens. Create them in the Web UI. A token needs `cxdb:read` to connect. Write tools also require `cxdb:write`. + +## Tools + +- `cxdb_list_contexts` +- `cxdb_search_contexts` +- `cxdb_get_context` +- `cxdb_get_turns` +- `cxdb_get_provenance` +- `cxdb_create_context` +- `cxdb_append_message` +- `cxdb_append_turn` + +Use `turn_id` with `cxdb_get_turns` to hydrate one complete turn. A bounded list response is a summary and can contain truncated strings. + +## Security boundary + +The gateway protects `/mcp` and gateway `/v1` routes. The direct Rust binary port 9009 and HTTP port 9010 keep their current behavior. Do not expose those direct ports when the gateway must be the authentication boundary. diff --git a/frontend/app/globals.css b/frontend/app/globals.css index f032506..a66569f 100644 --- a/frontend/app/globals.css +++ b/frontend/app/globals.css @@ -125,7 +125,10 @@ html, body { margin: 0; padding: 0; + width: 100%; min-height: 100vh; + min-height: 100dvh; + overflow-x: hidden; background: var(--bg); color: var(--text); font-family: 'Inter', system-ui, -apple-system, sans-serif; @@ -133,6 +136,23 @@ body { -moz-osx-font-smoothing: grayscale; } +button, +a, +input, +select, +textarea { + touch-action: manipulation; +} + +/* Prevent iOS from zooming the viewport when a form control receives focus. */ +@media (max-width: 767px) { + input, + select, + textarea { + font-size: 16px !important; + } +} + /* Custom scrollbar for dark theme */ ::-webkit-scrollbar { width: 8px; diff --git a/frontend/components/ContextDebugger.tsx b/frontend/components/ContextDebugger.tsx index 1645cf1..daddf83 100644 --- a/frontend/components/ContextDebugger.tsx +++ b/frontend/components/ContextDebugger.tsx @@ -1,10 +1,13 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + 'use client'; import { useEffect, useMemo, useRef, useState, useCallback } from 'react'; import type { Turn, TurnResponse, DebugEvent } from '@/types'; -import { Layers, Hash, X, Copy, Search, Loader2, AlertCircle, GitBranch, ChevronDown, ChevronRight, Terminal, MessageSquare, Wrench, CheckCircle, XCircle, Folder, Zap, Database } from './icons'; +import { Layers, Hash, X, Copy, Search, Loader2, AlertCircle, AlertTriangle, GitBranch, ChevronDown, ChevronRight, Terminal, MessageSquare, Wrench, CheckCircle, XCircle, Folder, Zap, Database } from './icons'; import { cn, trunc, safeStringify, formatTime, contentPreview } from '@/lib/utils'; -import { fetchTurns, fetchFsDirectory, ApiError } from '@/lib/api'; +import { fetchTurn, fetchTurns, fetchFsDirectory, ApiError } from '@/lib/api'; import { FileBrowser } from './FileBrowser'; import { FileViewer } from './FileViewer'; import { TryRenderCanonical, isConversationItem } from './ConversationRenderer'; @@ -17,6 +20,9 @@ import { useRendererManifest } from '@/lib/use-renderer'; import { getItemTypeLabel, getItemTypeColors } from '@/types/conversation'; import type { ConversationItem, ItemType } from '@/types/conversation'; +const TURN_PAGE_SIZE = 100; +const TURN_LIST_STRING_LIMIT = 512; + // View tabs for the right panel type DetailView = 'turn' | 'provenance'; @@ -149,6 +155,21 @@ function extractToolCalls(turn: Turn): Array<{ id: string; name: string; argumen })); } + // A legacy ToolCall is often stored as its own turn rather than inside an + // assistant message. Treat that one payload as a single call so its result + // can be matched by tool_call_id as well. + if (turn.declared_type?.type_id.includes('ToolCall')) { + const id = data.id ?? data.call_id ?? data['1']; + const name = data.name ?? data['2']; + if (id !== undefined && name !== undefined) { + return [{ + id: String(id), + name: String(name), + arguments: String(data.arguments ?? data.args ?? data['3'] ?? '{}'), + }]; + } + } + // Legacy extraction const toolCalls = data.tool_calls as Array> | undefined; if (!Array.isArray(toolCalls)) return []; @@ -161,6 +182,39 @@ function extractToolCalls(turn: Turn): Array<{ id: string; name: string; argumen })); } +/** + * A v2 assistant turn can carry its result in the same payload as the call. + * Such a result is already exact when this turn has been hydrated. Do not + * replace it with a lookup for a separate result turn. + */ +function hasEmbeddedToolResult(turn: Turn, toolCallId: string): boolean { + const data = turn.data as Record | undefined; + if (!data) return false; + + if (isConversationItem(data) && data.item_type === 'assistant_turn' && data.turn?.tool_calls) { + return data.turn.tool_calls.some(toolCall => { + if (toolCall.id !== toolCallId) return false; + return toolCall.result !== undefined + || toolCall.error !== undefined + || toolCall.streaming_output !== undefined; + }); + } + + // Keep compatibility with legacy payloads that use numeric msgpack keys. + const toolCalls = data.tool_calls as Array> | undefined; + if (!Array.isArray(toolCalls)) return false; + + return toolCalls.some(toolCall => { + const id = String(toolCall.id ?? toolCall['1'] ?? ''); + if (id !== toolCallId) return false; + return toolCall.result !== undefined + || toolCall.error !== undefined + || toolCall.streaming_output !== undefined + || toolCall['9'] !== undefined + || toolCall['10'] !== undefined; + }); +} + // Extract tool result info - handles canonical types and legacy formats function extractToolResult(turn: Turn): { toolCallId: string; content: string; isError: boolean } | null { const data = turn.data as Record | undefined; @@ -414,6 +468,128 @@ function TurnContentView({ turn }: { turn: Turn }) { return ; } +type ToolResultHydration = + | { state: 'loading' } + | { state: 'ready'; turn: Turn } + | { state: 'error' }; + +interface ToolResultMatchesProps { + contextId: string; + turn: Turn; + resultTurns: Map; +} + +/** + * Show results for calls which are represented by separate turns. + * + * The list endpoint is deliberately bounded for large traces. A result found + * in that list can therefore still contain a 512-character prefix. Hydrate + * each matched result by ID before rendering it. Missing means "not in the + * loaded page", not "the store has no such result". + */ +function ToolResultMatches({ contextId, turn, resultTurns }: ToolResultMatchesProps) { + const calls = useMemo(() => extractToolCalls(turn), [turn]); + const matches = useMemo(() => calls + .filter(call => !hasEmbeddedToolResult(turn, call.id)) + .map(call => ({ + call, + resultTurn: resultTurns.get(call.id) ?? null, + key: `${call.id}:${resultTurns.get(call.id)?.turn_id ?? 'missing'}`, + })), [calls, resultTurns, turn]); + const matchKey = matches.map(match => match.key).join('|'); + const [hydration, setHydration] = useState>({}); + + useEffect(() => { + let cancelled = false; + const initial: Record = {}; + for (const match of matches) { + if (match.resultTurn) initial[match.key] = { state: 'loading' }; + } + setHydration(initial); + + const matchedResults = matches.filter((match): match is typeof match & { resultTurn: Turn } => ( + match.resultTurn !== null + )); + if (matchedResults.length === 0) return () => { cancelled = true; }; + + for (const match of matchedResults) { + fetchTurn(contextId, match.resultTurn.turn_id) + .then(exactTurn => { + if (!cancelled) { + setHydration(previous => ({ + ...previous, + [match.key]: { state: 'ready', turn: exactTurn }, + })); + } + }) + .catch(() => { + if (!cancelled) { + setHydration(previous => ({ ...previous, [match.key]: { state: 'error' } })); + } + }); + } + + return () => { cancelled = true; }; + }, [contextId, matchKey, matches]); + + if (matches.length === 0) return null; + + return ( + {matches.length}} + > +
+ {matches.map(match => { + const result = hydration[match.key]; + return ( +
+
+ + {match.call.name} + {match.call.id} + {match.resultTurn && ( + + Turn #{match.resultTurn.turn_id} + + )} +
+ + {!match.resultTurn ? ( +
+ + No separate result turn in the loaded page. +
+ ) : result?.state === 'error' ? ( +
+ + Failed to load the complete result turn. +
+ ) : result?.state === 'ready' ? ( + + ) : ( +
+ + Loading the complete result… +
+ )} +
+ ); + })} +
+
+ ); +} + // Format arguments JSON for display function formatArguments(args: string): string { try { @@ -437,11 +613,10 @@ interface ContextDebuggerProps { onNavigateToContext?: (contextId: string) => void; } -const TURNS_PAGE_SIZE = 100; - export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initialTurnId, onTurnChange, onNavigateToContext }: ContextDebuggerProps) { const containerRef = useRef(null); const turnListRef = useRef(null); + const lastResetContextIdRef = useRef(null); const [query, setQuery] = useState(''); const [selectedIdx, setSelectedIdx] = useState(0); const [copied, setCopied] = useState<'context' | 'event' | null>(null); @@ -449,9 +624,15 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial // Data fetching state const [loading, setLoading] = useState(false); - const [loadingMore, setLoadingMore] = useState(false); + const [loadingOlder, setLoadingOlder] = useState(false); const [error, setError] = useState(null); const [data, setData] = useState(null); + const [hasMoreTurns, setHasMoreTurns] = useState(false); + const [selectedTurnDetail, setSelectedTurnDetail] = useState(null); + const [detailLoading, setDetailLoading] = useState(false); + const [detailError, setDetailError] = useState(null); + const [searchHydrating, setSearchHydrating] = useState(false); + const [searchHydrationError, setSearchHydrationError] = useState(null); // Live observer state const [newTurnIds, setNewTurnIds] = useState>(new Set()); @@ -477,11 +658,13 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial try { const response = await fetchTurns(contextId, { - limit: TURNS_PAGE_SIZE, + limit: TURN_PAGE_SIZE, view: 'typed', include_unknown: true, + string_limit: TURN_LIST_STRING_LIMIT, }); setData(response); + setHasMoreTurns(response.turns.length === TURN_PAGE_SIZE); } catch (err) { if (err instanceof ApiError) { setError(err.message); @@ -494,31 +677,52 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial } }, [contextId]); - // Load older turns using pagination cursor - const loadMore = useCallback(async () => { - if (!contextId || !data?.next_before_turn_id) return; + const loadOlderTurns = useCallback(async () => { + if (!contextId || !data?.next_before_turn_id || loadingOlder) return; + + setLoadingOlder(true); + setError(null); - setLoadingMore(true); try { const response = await fetchTurns(contextId, { - limit: TURNS_PAGE_SIZE, + limit: TURN_PAGE_SIZE, before_turn_id: data.next_before_turn_id, view: 'typed', include_unknown: true, + string_limit: TURN_LIST_STRING_LIMIT, }); - const prepended = response.turns.length; - setData(prev => prev ? { - ...prev, - turns: [...response.turns, ...prev.turns], - next_before_turn_id: response.next_before_turn_id, - } : response); - setSelectedIdx(prev => prev + prepended); - } catch { - // Keep existing data on failure + const seen = new Set(data.turns.map(turn => turn.turn_id)); + const olderTurns = response.turns.filter(turn => !seen.has(turn.turn_id)); + setData(prev => { + if (!prev) return response; + const currentIds = new Set(prev.turns.map(turn => turn.turn_id)); + const turnsToAdd = olderTurns.filter(turn => !currentIds.has(turn.turn_id)); + if (turnsToAdd.length === 0) { + return { + ...prev, + next_before_turn_id: response.next_before_turn_id, + }; + } + return { + ...prev, + turns: [...turnsToAdd, ...prev.turns], + next_before_turn_id: response.next_before_turn_id, + }; + }); + if (olderTurns.length > 0) { + setSelectedIdx(idx => idx + olderTurns.length); + } + setHasMoreTurns(response.turns.length === TURN_PAGE_SIZE); + } catch (err) { + if (err instanceof ApiError) { + setError(err.message); + } else { + setError('Failed to fetch older turns'); + } } finally { - setLoadingMore(false); + setLoadingOlder(false); } - }, [contextId, data?.next_before_turn_id]); + }, [contextId, data?.next_before_turn_id, data?.turns, loadingOlder]); useEffect(() => { if (isOpen && contextId) { @@ -568,10 +772,14 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial }, []); // Filter turns by search query + const hasSearchQuery = query.trim().length > 0; + const isSummaryPage = typeof data?.meta.string_limit === 'number'; const filteredTurns = useMemo(() => { if (!data?.turns) return []; const q = query.trim().toLowerCase(); if (!q) return data.turns; + // Never present prefix-only filtering as a complete search result. + if (typeof data.meta.string_limit === 'number') return []; return data.turns.filter(turn => { const content = extractContent(turn)?.toLowerCase() ?? ''; @@ -583,7 +791,118 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial }, [data, query]); // Selected turn - const selectedTurn = filteredTurns[selectedIdx] ?? null; + const selectedListTurn = filteredTurns[selectedIdx] ?? null; + const selectedTurn = selectedTurnDetail?.turn_id === selectedListTurn?.turn_id + ? selectedTurnDetail + : selectedListTurn; + const selectedTurnIsExact = !isSummaryPage + || selectedTurnDetail?.turn_id === selectedListTurn?.turn_id; + + // Build a bounded index from the turns already loaded. Matching must not + // fetch the complete context history, especially for large traces. + const separateToolResultTurns = useMemo(() => { + const resultTurns = new Map(); + for (const turn of data?.turns ?? []) { + const result = extractToolResult(turn); + if (result && result.toolCallId !== 'unknown') { + resultTurns.set(result.toolCallId, turn); + } + } + return resultTurns; + }, [data]); + + // List pages carry bounded string prefixes. Fetch only the selected turn's + // complete payload for the detail renderer. + useEffect(() => { + const turnId = selectedListTurn?.turn_id; + if (!turnId || typeof data?.meta.string_limit !== 'number') { + setSelectedTurnDetail(null); + setDetailLoading(false); + setDetailError(null); + return; + } + + let cancelled = false; + let requestStarted = false; + setSelectedTurnDetail(null); + setDetailLoading(true); + setDetailError(null); + // Selection can move from the first rendered row to the followed tail in + // the same render cycle. Debouncing avoids transferring the discarded + // detail and also coalesces rapid keyboard navigation. + const timer = window.setTimeout(() => { + requestStarted = true; + fetchTurn(contextId, turnId) + .then(turn => { + if (!cancelled) setSelectedTurnDetail(turn); + }) + .catch(() => { + if (!cancelled) setDetailError('Failed to load the complete turn.'); + }) + .finally(() => { + if (!cancelled) setDetailLoading(false); + }); + }, 25); + return () => { + cancelled = true; + if (!requestStarted) window.clearTimeout(timer); + }; + }, [contextId, data?.meta.string_limit, selectedListTurn?.turn_id]); + + // Preserve full-text filtering semantics: the common browsing path uses + // summaries, while entering a query hydrates every currently loaded turn. + useEffect(() => { + if ( + !hasSearchQuery + || !data + || typeof data.meta.string_limit !== 'number' + ) { + if (!hasSearchQuery) { + setSearchHydrating(false); + setSearchHydrationError(null); + } + return; + } + let cancelled = false; + setSearchHydrating(true); + setSearchHydrationError(null); + const loadedTurnIds = data.turns.map(turn => turn.turn_id); + const hydrateLoadedTurns = async () => { + const hydrated = new Map(); + const concurrency = 16; + for (let offset = 0; offset < loadedTurnIds.length; offset += concurrency) { + const page = await Promise.all( + loadedTurnIds.slice(offset, offset + concurrency).map(turnId => fetchTurn(contextId, turnId)) + ); + for (const turn of page) hydrated.set(turn.turn_id, turn); + } + return hydrated; + }; + hydrateLoadedTurns() + .then(hydrated => { + if (!cancelled) { + setData(prev => { + if (!prev) return prev; + const { string_limit: _stringLimit, ...completeMeta } = prev.meta; + return { + ...prev, + meta: completeMeta, + turns: prev.turns.map(turn => hydrated.get(turn.turn_id) ?? turn), + }; + }); + setSelectedTurnDetail(null); + } + }) + .catch(() => { + if (!cancelled) { + setSearchHydrationError('Failed to load complete turns for search.'); + } + }) + .finally(() => { + if (!cancelled) setSearchHydrating(false); + }); + return () => { cancelled = true; }; + }, [contextId, data, hasSearchQuery]); // Detect filesystem for selected turn const selectedTurnId = selectedTurn?.turn_id; @@ -628,12 +947,42 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial if (!data?.turns || initialTurnApplied) return; if (initialTurnId) { - // Find and select the specified turn const idx = filteredTurns.findIndex(t => t.turn_id === initialTurnId); if (idx >= 0) { setSelectedIdx(idx); setInitialTurnApplied(true); + return; } + + // A deep link can target a turn outside the bounded first page. Hydrate + // that exact turn and add it to the list so the URL and visible selection + // cannot disagree. + let cancelled = false; + setDetailLoading(true); + setDetailError(null); + fetchTurn(contextId, initialTurnId) + .then(turn => { + if (cancelled) return; + setData(previous => { + if (!previous || previous.turns.some(item => item.turn_id === turn.turn_id)) { + return previous; + } + return { ...previous, turns: [turn, ...previous.turns] }; + }); + setSelectedIdx(0); + setSelectedTurnDetail(turn); + setInitialTurnApplied(true); + }) + .catch(() => { + if (!cancelled) { + setDetailError('Failed to load the linked turn.'); + setInitialTurnApplied(true); + } + }) + .finally(() => { + if (!cancelled) setDetailLoading(false); + }); + return () => { cancelled = true; }; } else if (filteredTurns.length > 0) { // No initial turn specified, notify parent of first turn const firstTurn = filteredTurns[0]; @@ -642,11 +991,11 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial } setInitialTurnApplied(true); } - }, [data, initialTurnId, initialTurnApplied, filteredTurns, onTurnChange]); + }, [contextId, data, initialTurnId, initialTurnApplied, filteredTurns, onTurnChange]); // Count stats - count both tool_call turns AND tool_calls embedded in assistant turns const stats = useMemo(() => { - if (!data?.turns) return { total: 0, loaded: 0, toolCalls: 0, errors: 0 }; + if (!data?.turns) return { loaded: 0, total: 0, toolCalls: 0, errors: 0 }; let toolCalls = 0; let errors = 0; for (const turn of data.turns) { @@ -661,9 +1010,12 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial const result = extractToolResult(turn); if (result?.isError) errors++; } - const headId = data.meta?.head_turn_id; - const total = headId && headId !== '0' ? (data.meta?.head_depth ?? 0) + 1 : 0; - return { total, loaded: data.turns.length, toolCalls, errors }; + return { + loaded: data.turns.length, + total: data.meta.head_turn_id === '0' ? 0 : data.meta.head_depth + 1, + toolCalls, + errors, + }; }, [data]); // Auto-select last turn when following and new turns arrive @@ -696,17 +1048,33 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial // Reset state when modal opens/closes useEffect(() => { - if (!isOpen) return; + if (!isOpen) { + lastResetContextIdRef.current = null; + return; + } + + // Don't reset state on URL turn-id changes while open; only reset on open/context changes. + if (lastResetContextIdRef.current === contextId) { + return; + } + lastResetContextIdRef.current = contextId; + setQuery(''); // Only reset to 0 if no initialTurnId; otherwise let the initialTurn effect handle it if (!initialTurnId) { setSelectedIdx(0); } - setInitialTurnApplied(false); setCopied(null); requestAnimationFrame(() => containerRef.current?.focus()); }, [isOpen, contextId, initialTurnId]); + // Allow URL-driven turn selection changes (e.g. browser back/forward) to re-apply without + // wiping the user's current filter query. + useEffect(() => { + if (!isOpen) return; + setInitialTurnApplied(false); + }, [isOpen, initialTurnId]); + // Clear copied state after delay useEffect(() => { if (!copied) return; @@ -718,9 +1086,30 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial const handleCopy = async (kind: 'context' | 'event') => { try { - const text = kind === 'context' - ? safeStringify(data ?? { error: 'No data' }) - : safeStringify(selectedTurn ?? {}); + let value: unknown; + if (kind === 'context' && data && typeof data.meta.string_limit === 'number') { + const complete = await fetchTurns(contextId, { + limit: data.turns.length, + view: 'typed', + include_unknown: true, + }); + setData(complete); + value = complete; + } else if ( + kind === 'event' + && selectedListTurn + && typeof data?.meta.string_limit === 'number' + && selectedTurn?.turn_id !== selectedTurnDetail?.turn_id + ) { + const complete = await fetchTurn(contextId, selectedListTurn.turn_id); + setSelectedTurnDetail(complete); + value = complete; + } else { + value = kind === 'context' + ? data ?? { error: 'No data' } + : selectedTurn ?? {}; + } + const text = safeStringify(value); await navigator.clipboard.writeText(text); setCopied(kind); } catch { @@ -778,51 +1167,56 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial ref={containerRef} tabIndex={-1} onKeyDown={handleKeyDown} - className="h-full w-full outline-none" + className="flex h-[100dvh] w-full flex-col outline-none" data-context-debugger > {/* Header - more compact */} -
-
-
- - Context {contextId} +
+
+
+ + Context {contextId}
{data && ( -
- {stats.loaded < stats.total ? `${stats.loaded} of ${stats.total} turns` : `${stats.total} turns`} - {stats.toolCalls} tool calls +
+ {stats.loaded} of {stats.total} turns loaded + {stats.toolCalls} tool calls {stats.errors > 0 && ( - {stats.errors} errors + {stats.errors} errors )}
)}
-
+
@@ -830,9 +1224,9 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial
{/* Body */} -
+
{/* Left: Turn list - more compact */} -
+
@@ -852,6 +1246,24 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial className={cn('overflow-y-auto relative', hasFilesystem ? 'flex-1 min-h-0' : 'flex-1')} data-debug-event-list > + {!loading && !error && data && hasMoreTurns && !query.trim() && ( +
+ +
+ )} {loading ? (
@@ -862,22 +1274,22 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial {error}
+ ) : searchHydrating ? ( +
+ + Loading complete turns for search… +
+ ) : searchHydrationError ? ( +
+ + {searchHydrationError} +
) : filteredTurns.length === 0 ? (
{data?.turns.length === 0 ? 'No turns.' : 'No matches.'}
) : ( - <> - {data && data.turns.length > 0 && data.turns[0].depth > 0 && ( - - )} - {filteredTurns.map((turn, idx) => { + filteredTurns.map((turn, idx) => { const kind = detectTurnKind(turn); const colors = getKindColors(kind); const isSelected = idx === selectedIdx; @@ -921,8 +1333,7 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial
); - })} - + }) )} {/* Resume following indicator */} @@ -957,7 +1368,7 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial
{/* Right: Detail view */} -
+
{/* File viewer overlay */} {selectedFilePath && selectedTurn && ( {/* Detail view tabs */} -
+
)}
{/* Turn header (when viewing turn) */} {detailView === 'turn' && ( -
+
{getKindLabel(detectTurnKind(selectedTurn))}
- + Turn #{selectedTurn.turn_id} • Depth {selectedTurn.depth}
@@ -1031,48 +1444,84 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial {/* Content area - Turn view */} {detailView === 'turn' && ( -
- {/* Primary content view - uses dynamic renderer registry */} - - - {/* Collapsible metadata */} - - {selectedTurn.declared_type?.type_id?.split('.').pop()} - - } - > -
-
Turn ID
-
{selectedTurn.turn_id}
-
Parent
-
{selectedTurn.parent_turn_id || '(root)'}
-
Depth
-
{selectedTurn.depth}
- {selectedTurn.declared_type && ( - <> -
Type
-
- {selectedTurn.declared_type.type_id}@{selectedTurn.declared_type.type_version} -
- +
+ {!selectedTurnIsExact ? ( +
+ {detailError ? ( + + ) : ( + )} + {detailError ?? (detailLoading + ? 'Loading full turn…' + : 'Waiting for full turn…')}
- - - {/* Collapsible raw payload */} - -
-                        {safeStringify(selectedTurn.data)}
-                      
-
+ ) : ( + <> + {selectedTurn.projection_error && ( +
+ +
+
Typed projection unavailable
+
+ {selectedTurn.projection_error.message} +
+
+
+ )} + + {/* Primary content view - uses dynamic renderer registry */} + + + + + {/* Collapsible metadata */} + + {selectedTurn.declared_type?.type_id?.split('.').pop()} + + } + > +
+
Turn ID
+
{selectedTurn.turn_id}
+
Parent
+
{selectedTurn.parent_turn_id || '(root)'}
+
Depth
+
{selectedTurn.depth}
+ {selectedTurn.declared_type && ( + <> +
Type
+
+ {selectedTurn.declared_type.type_id}@{selectedTurn.declared_type.type_version} +
+ + )} +
+
+ + {/* Collapsible raw payload */} + +
+                            {safeStringify(selectedTurn.data ?? selectedTurn.projection_error)}
+                          
+
+ + )}
)} @@ -1093,7 +1542,7 @@ export function ContextDebugger({ contextId, isOpen, onClose, lastEvent, initial )} {/* Footer */} -
+
j/k Navigate F Follow ⌘K Search diff --git a/frontend/components/CxdbApp.tsx b/frontend/components/CxdbApp.tsx index 5a27d6b..771c72f 100644 --- a/frontend/components/CxdbApp.tsx +++ b/frontend/components/CxdbApp.tsx @@ -1,18 +1,22 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + 'use client'; import { useState, useCallback, useEffect, useMemo, useRef } from 'react'; import { ContextDebugger } from '@/components/ContextDebugger'; import { ContextList } from '@/components/ContextList'; import type { ContextEntry, StoreEvent } from '@/types'; -import { Database, Layers, Plus, X, AlertCircle, Check, Zap, Radio, ChevronDown, Filter } from '@/components/icons'; +import { Database, Layers, Plus, X, AlertCircle, Check, Radio, ChevronDown, Filter, Lock } from '@/components/icons'; import { ThemeSelector } from '@/components/ThemeSelector'; import { getTagColor } from '@/lib/clientTags'; import { cn, normalizeContextId } from '@/lib/utils'; import { healthCheck, fetchContexts, searchContexts } from '@/lib/api'; import { validate as validateCql, buildFallbackQuery, appendSearchCriterionClause, extractSearchCriteriaClauses } from '@/lib/cql'; import { useEventStream, useMockEventGenerator, useUrlRouter, parseUrl, type RouteState } from '@/hooks'; -import { ConnectionStatus, ActivityFeed } from '@/components/live'; +import { ActivityFeed } from '@/components/live'; import { ServerHealthDashboard } from '@/components/dashboard'; +import { TokenManagement } from '@/components/TokenManagement'; export default function CxdbApp() { const [contexts, setContexts] = useState([]); @@ -37,6 +41,7 @@ export default function CxdbApp() { // Environment filter state const [selectedEnv, setSelectedEnv] = useState<'all' | 'prod' | 'stage' | 'dev'>('all'); + const [tokenManagementOpen, setTokenManagementOpen] = useState(false); // URL routing - parse URL on mount and handle changes const handleRouteChange = useCallback((state: RouteState) => { @@ -186,6 +191,14 @@ export default function CxdbApp() { // Mock event generator for demo const { startMockEvents, stopMockEvents } = useMockEventGenerator(mockEmit); + // One demo control owns the generator lifecycle. This also makes the demo + // useful on touch devices without a second hidden action. + useEffect(() => { + if (!mockMode) return; + startMockEvents(2000); + return stopMockEvents; + }, [mockMode, startMockEvents, stopMockEvents]); + // Fetch contexts helper const fetchContextsData = useCallback(async () => { try { @@ -424,7 +437,7 @@ export default function CxdbApp() { if (e.metaKey || e.ctrlKey || e.altKey) return; // Only handle j/k/o when debugger is closed and viewing contexts (not activity) - if (!debuggerOpen && !showActivityFeed) { + if (!debuggerOpen && !showActivityFeed && !tokenManagementOpen) { if (e.key === 'j' || e.key === 'ArrowDown') { e.preventDefault(); setFocusedContextIndex(prev => @@ -454,35 +467,35 @@ export default function CxdbApp() { }; window.addEventListener('keydown', handleKey); return () => window.removeEventListener('keydown', handleKey); - }, [debuggerOpen, showActivityFeed, filteredContexts, focusedContextIndex, handleSelectContext]); + }, [debuggerOpen, showActivityFeed, tokenManagementOpen, filteredContexts, focusedContextIndex, handleSelectContext]); return ( -
+
{/* Header */} -
+
{/* Left: Logo + Title + Env Pills */} -
+
-
+

CXDB

-

AI Context Store

+

AI Context Store

{/* Environment Filter Pills - vertically centered with logo */} -
+
{(['all', 'prod', 'stage', 'dev'] as const).map((env) => (
{/* Right: Controls */} -
- {/* Theme selector */} +
- {/* Mock mode toggle */} + + - {/* Demo button (mock mode only) */} - {mockMode && ( - - )} - - {/* Connection status */} - - - {/* Server status indicator */}
+ )} role="status" aria-label={serverStatus === 'online' ? 'Server online' : serverStatus === 'offline' ? 'Server offline' : 'Checking server status'}> - {serverStatus === 'online' ? 'Server online' : - serverStatus === 'offline' ? 'Server offline' : - 'Checking...'} + {serverStatus === 'online' ? 'Server online' : serverStatus === 'offline' ? 'Server offline' : 'Checking...'}
{/* Main content */} -
+ ); } diff --git a/frontend/components/QuestRenderer.tsx b/frontend/components/QuestRenderer.tsx index a4633e2..ffba37c 100644 --- a/frontend/components/QuestRenderer.tsx +++ b/frontend/components/QuestRenderer.tsx @@ -1,3 +1,6 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + 'use client'; import { useState } from 'react'; @@ -157,7 +160,7 @@ function ActionSummaryCard({ summary }: { summary: ActionSummary }) { const skipped = Number(summary.skipped) || 0; return ( -
+
{total}
Total
diff --git a/frontend/components/ThemeSelector.tsx b/frontend/components/ThemeSelector.tsx index ab8b0ff..27eee84 100644 --- a/frontend/components/ThemeSelector.tsx +++ b/frontend/components/ThemeSelector.tsx @@ -1,3 +1,6 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + 'use client'; import { useState, useRef, useEffect } from 'react'; @@ -73,6 +76,7 @@ export function ThemeSelector({ className }: ThemeSelectorProps) { )} aria-expanded={isOpen} aria-haspopup="listbox" + aria-label={`Theme: ${theme.name}`} > @@ -88,7 +92,7 @@ export function ThemeSelector({ className }: ThemeSelectorProps) { {isOpen && (
void; +} + +function formatDate(value?: string | null): string { + if (!value) return 'Never'; + const date = new Date(value); + if (Number.isNaN(date.getTime())) return 'Unknown'; + return new Intl.DateTimeFormat(undefined, { dateStyle: 'medium', timeStyle: 'short' }).format(date); +} + +function isExpired(value: string): boolean { + const date = new Date(value); + return !Number.isNaN(date.getTime()) && date.getTime() <= Date.now(); +} + +export function TokenManagement({ isOpen, onClose }: TokenManagementProps) { + const [tokens, setTokens] = useState([]); + const [csrfToken, setCsrfToken] = useState(null); + const [loading, setLoading] = useState(false); + const [creating, setCreating] = useState(false); + const [error, setError] = useState(null); + const [name, setName] = useState(''); + const [includeWrite, setIncludeWrite] = useState(false); + const [expiresAt, setExpiresAt] = useState(''); + const [newPlaintext, setNewPlaintext] = useState(null); + const [copyStatus, setCopyStatus] = useState<'idle' | 'copied' | 'failed'>('idle'); + const [revokingId, setRevokingId] = useState(null); + const [confirmingId, setConfirmingId] = useState(null); + + useEffect(() => { + if (!isOpen) { + setNewPlaintext(null); + setCopyStatus('idle'); + setConfirmingId(null); + return; + } + + let cancelled = false; + setLoading(true); + setError(null); + const load = async () => { + try { + const user = await fetchCurrentUser(); + const listedTokens = await fetchAPITokens(); + if (!cancelled) { + setCsrfToken(user.csrf_token); + setTokens(listedTokens); + } + } catch (err) { + if (!cancelled) setError(err instanceof Error ? err.message : 'Unable to load API tokens.'); + } finally { + if (!cancelled) setLoading(false); + } + }; + void load(); + return () => { cancelled = true; }; + }, [isOpen]); + + if (!isOpen) return null; + + const closePanel = () => { + setNewPlaintext(null); + setCopyStatus('idle'); + onClose(); + }; + + const handleCreate = async (event: React.FormEvent) => { + event.preventDefault(); + const trimmedName = name.trim(); + if (!trimmedName) { + setError('Enter a name for this token.'); + return; + } + if (!csrfToken) { + setError('Your browser session is not ready. Reload and try again.'); + return; + } + + let expires: string | undefined; + if (expiresAt) { + const date = new Date(expiresAt); + if (Number.isNaN(date.getTime()) || date.getTime() <= Date.now()) { + setError('Expiry must be a future date.'); + return; + } + expires = date.toISOString(); + } + + setCreating(true); + setError(null); + try { + const result = await createAPIToken(csrfToken, { + name: trimmedName, + scopes: includeWrite ? ['cxdb:read', 'cxdb:write'] : ['cxdb:read'], + ...(expires ? { expires_at: expires } : {}), + }); + setTokens(previous => [result.token, ...previous.filter(token => token.id !== result.token.id)]); + setNewPlaintext(result.plaintext); + setCopyStatus('idle'); + setName(''); + setIncludeWrite(false); + setExpiresAt(''); + } catch (err) { + setError(err instanceof Error ? err.message : 'Unable to create API token.'); + } finally { + setCreating(false); + } + }; + + const handleCopy = async () => { + if (!newPlaintext) return; + try { + await navigator.clipboard.writeText(newPlaintext); + setCopyStatus('copied'); + } catch { + setCopyStatus('failed'); + } + }; + + const handleRevoke = async (tokenId: string) => { + if (!csrfToken) { + setError('Your browser session is not ready. Reload and try again.'); + return; + } + setRevokingId(tokenId); + setError(null); + try { + await revokeAPIToken(csrfToken, tokenId); + setTokens(previous => previous.map(token => token.id === tokenId + ? { ...token, revoked_at: new Date().toISOString() } + : token)); + setConfirmingId(null); + } catch (err) { + setError(err instanceof Error ? err.message : 'Unable to revoke API token.'); + } finally { + setRevokingId(null); + } + }; + + return ( +
+
+
+
+
+
+

API tokens

+

Manage personal access for tools and scripts.

+
+
+ +
+ +
+ {error &&
{error}
} + + {newPlaintext && ( +
+

Token created

+

This secret is shown only once. Copy it now. It will not be available again.

+
+ + +
+ {copyStatus === 'failed' &&

Copy failed. Select the secret and copy it manually.

} +
+ )} + +
+

Create token

+
+ +
Scopes + + +
+ +
+
+
+ +
+

Your tokens

+ {loading ?
Loading tokens...
: tokens.length === 0 ?
No API tokens yet.
: ( +
+ {tokens.map(token => { + const revoked = Boolean(token.revoked_at); + const expired = isExpired(token.expires_at); + return
+
+

{token.name}

{revoked && Revoked}{!revoked && expired && Expired}
+

{token.prefix}

+
{token.scopes.map(scope => {scope})}
+
Created:
{formatDate(token.created_at)}
Expires:
{formatDate(token.expires_at)}
Last used:
{formatDate(token.last_used_at)}
+
+ {!revoked && (confirmingId === token.id ?
Revoke this token?
: )} +
+
; + })} +
+ )} +
+
+
+
+ ); +} diff --git a/frontend/components/dashboard/CapacityGauge.tsx b/frontend/components/dashboard/CapacityGauge.tsx index c80c8f0..a8db1e1 100644 --- a/frontend/components/dashboard/CapacityGauge.tsx +++ b/frontend/components/dashboard/CapacityGauge.tsx @@ -1,3 +1,6 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + 'use client'; import { cn } from '@/lib/utils'; @@ -79,9 +82,9 @@ export function CapacityGauge({ {/* Main horizontal gauge */}
-
+

In-Memory Index Capacity

-
+
{(capacityRatio * 100).toFixed(0)}% diff --git a/frontend/components/dashboard/ObjectCountsCard.tsx b/frontend/components/dashboard/ObjectCountsCard.tsx index e3d0115..9f99898 100644 --- a/frontend/components/dashboard/ObjectCountsCard.tsx +++ b/frontend/components/dashboard/ObjectCountsCard.tsx @@ -1,3 +1,6 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + 'use client'; import { cn } from '@/lib/utils'; @@ -29,21 +32,21 @@ export function ObjectCountsCard({ objects, previousObjects, filesystem, classNa
{/* Main counts */} -
+
-
+
{formatCount(objects.contexts_total)}
contexts
-
+
{formatCount(objects.turns_total)}
turns
-
+
{formatCount(objects.blobs_total)}
blobs
diff --git a/frontend/components/dashboard/ServerHealthDashboard.tsx b/frontend/components/dashboard/ServerHealthDashboard.tsx index a1b0b0f..abfb6bd 100644 --- a/frontend/components/dashboard/ServerHealthDashboard.tsx +++ b/frontend/components/dashboard/ServerHealthDashboard.tsx @@ -1,3 +1,6 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + 'use client'; import { RefreshCw, WifiOff } from '@/components/icons'; @@ -20,8 +23,8 @@ function DashboardSkeleton() { return (
{/* Gauge skeleton */} -
-
+
+
{/* Cards skeleton */} @@ -56,7 +59,7 @@ export function ServerHealthDashboard({ // Loading state if (status === 'loading' && !data) { return ( -
+
); @@ -65,8 +68,8 @@ export function ServerHealthDashboard({ // Offline state (no data at all) if (isOffline) { return ( -
-
+
+

Server Offline

@@ -93,10 +96,10 @@ export function ServerHealthDashboard({ if (!data) return null; return ( -

+
{/* Stale warning */} {isStale && ( -
+
Data may be stale · Last updated: {lastUpdated} diff --git a/frontend/components/dashboard/SessionsErrorsBar.tsx b/frontend/components/dashboard/SessionsErrorsBar.tsx index daab336..cd5794d 100644 --- a/frontend/components/dashboard/SessionsErrorsBar.tsx +++ b/frontend/components/dashboard/SessionsErrorsBar.tsx @@ -1,3 +1,6 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + 'use client'; import { useState, useCallback } from 'react'; @@ -54,8 +57,8 @@ export function SessionsErrorsBar({ sessions, errors, perf, className }: Session const hasErrors = errors.total > 0; return ( -
-
+
+
{/* Sessions info */}
@@ -73,7 +76,7 @@ export function SessionsErrorsBar({ sessions, errors, perf, className }: Session
{/* Errors info - clickable when errors exist */} -
+
+`)) + +// VerifyOAuthTokenWithContext is useful to adapters that need request context. +func (s *OAuthServer) VerifyOAuthTokenWithContext(_ context.Context, token string) (*Session, error) { + return s.Verify(token) +} diff --git a/gateway/pkg/auth/oauth_server_test.go b/gateway/pkg/auth/oauth_server_test.go new file mode 100644 index 0000000..704bb11 --- /dev/null +++ b/gateway/pkg/auth/oauth_server_test.go @@ -0,0 +1,217 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + +package auth + +import ( + "context" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "net/url" + "path/filepath" + "strings" + "testing" + "time" +) + +func testOAuthServer(t *testing.T) (*OAuthServer, *SessionStore) { + t.Helper() + store, err := NewSessionStore(filepath.Join(t.TempDir(), "oauth.sqlite"), "session", time.Hour, "", false, "oauth-test-secret") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = store.Close() }) + server, err := NewOAuthServer(store, "https://cxdb.example", "/auth/login") + if err != nil { + t.Fatal(err) + } + return server, store +} + +func TestOAuthAuthorizationCodePKCEAndSingleUse(t *testing.T) { + server, store := testOAuthServer(t) + clientID, redirectURI := registerTestClient(t, server, "http://127.0.0.1:49152/callback") + sessionID, err := store.CreateForIdentity(context.Background(), "https://issuer.example", "alice", "alice@example.com", "Alice", "", "oidc", []string{"cxdb:read", "cxdb:write"}) + if err != nil { + t.Fatal(err) + } + verifier := strings.Repeat("a", 64) + digest := sha256.Sum256([]byte(verifier)) + request := oauthAuthorizationRequest{ + ClientID: clientID, RedirectURI: redirectURI, State: "client-state", + Challenge: base64.RawURLEncoding.EncodeToString(digest[:]), Scopes: []string{"cxdb:read", "cxdb:write"}, + Resource: server.resource, ExpiresAt: time.Now().Add(time.Minute).Unix(), + } + signed, err := server.signAuthorizationRequest(request) + if err != nil { + t.Fatal(err) + } + form := url.Values{"request": {signed}, "decision": {"allow"}} + authorizeRequest := httptest.NewRequest(http.MethodPost, "/oauth/authorize", strings.NewReader(form.Encode())) + authorizeRequest.Header.Set("Content-Type", "application/x-www-form-urlencoded") + setSessionCookie(t, store, sessionID, authorizeRequest) + authorizeResponse := httptest.NewRecorder() + server.AuthorizeHandler(authorizeResponse, authorizeRequest) + if authorizeResponse.Code != http.StatusFound { + t.Fatalf("authorize status = %d, body=%s", authorizeResponse.Code, authorizeResponse.Body.String()) + } + redirect, err := url.Parse(authorizeResponse.Header().Get("Location")) + if err != nil { + t.Fatal(err) + } + if redirect.Query().Get("state") != "client-state" || redirect.Query().Get("iss") != "https://cxdb.example" { + t.Fatalf("authorization response parameters = %s", redirect.RawQuery) + } + code := redirect.Query().Get("code") + tokenForm := url.Values{"grant_type": {"authorization_code"}, "code": {code}, "client_id": {clientID}, "redirect_uri": {redirectURI}, "code_verifier": {verifier}} + tokenRequest := httptest.NewRequest(http.MethodPost, "/oauth/token", strings.NewReader(tokenForm.Encode())) + tokenRequest.Header.Set("Content-Type", "application/x-www-form-urlencoded") + tokenResponse := httptest.NewRecorder() + server.TokenHandler(tokenResponse, tokenRequest) + if tokenResponse.Code != http.StatusOK { + t.Fatalf("token status = %d, body=%s", tokenResponse.Code, tokenResponse.Body.String()) + } + var tokenPayload map[string]any + if err := json.Unmarshal(tokenResponse.Body.Bytes(), &tokenPayload); err != nil { + t.Fatal(err) + } + verified, err := server.Verify(tokenPayload["access_token"].(string)) + if err != nil || !verified.HasScope("cxdb:write") || verified.Subject != "alice" { + t.Fatalf("verified token = %+v, err=%v", verified, err) + } + + replayRequest := httptest.NewRequest(http.MethodPost, "/oauth/token", strings.NewReader(tokenForm.Encode())) + replayRequest.Header.Set("Content-Type", "application/x-www-form-urlencoded") + replayResponse := httptest.NewRecorder() + server.TokenHandler(replayResponse, replayRequest) + if replayResponse.Code != http.StatusBadRequest || !strings.Contains(replayResponse.Body.String(), "invalid_grant") { + t.Fatalf("code replay accepted: status=%d body=%s", replayResponse.Code, replayResponse.Body.String()) + } + if _, err := server.Verify(tokenPayload["access_token"].(string)); err == nil { + t.Fatal("authorization-code replay did not revoke the issued access token") + } +} + +func TestOAuthRegistrationRejectsUnsafeRedirects(t *testing.T) { + server, _ := testOAuthServer(t) + for _, redirect := range []string{"http://example.com/callback", "https://user@example.com/callback", "https://example.com/callback#fragment", "file:///tmp/callback"} { + body := `{"redirect_uris":[` + strconvQuote(redirect) + `]}` + request := httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body)) + response := httptest.NewRecorder() + server.RegisterHandler(response, request) + if response.Code != http.StatusBadRequest { + t.Errorf("redirect %q status = %d", redirect, response.Code) + } + } +} + +func TestOAuthConsentCSPAllowsOnlyRegisteredCallbackOrigin(t *testing.T) { + server, store := testOAuthServer(t) + clientID, redirectURI := registerTestClient(t, server, "http://127.0.0.1:49152/callback") + sessionID, err := store.CreateForIdentity(context.Background(), "https://issuer.example", "alice", "alice@example.com", "Alice", "", "oidc", []string{"cxdb:read", "cxdb:write"}) + if err != nil { + t.Fatal(err) + } + challenge := base64.RawURLEncoding.EncodeToString(make([]byte, sha256.Size)) + query := url.Values{ + "response_type": {"code"}, "client_id": {clientID}, "redirect_uri": {redirectURI}, + "state": {"state"}, "scope": {"cxdb:read"}, "resource": {server.resource}, + "code_challenge": {challenge}, "code_challenge_method": {"S256"}, + } + request := httptest.NewRequest(http.MethodGet, "/oauth/authorize?"+query.Encode(), nil) + setSessionCookie(t, store, sessionID, request) + response := httptest.NewRecorder() + server.AuthorizeHandler(response, request) + if response.Code != http.StatusOK { + t.Fatalf("consent status = %d, body=%s", response.Code, response.Body.String()) + } + want := "default-src 'none'; form-action 'self' http://127.0.0.1:49152; frame-ancestors 'none'; base-uri 'none'" + if got := response.Header().Get("Content-Security-Policy"); got != want { + t.Fatalf("consent CSP = %q, want %q", got, want) + } +} + +func TestOAuthConsentCSPCallbackOrigins(t *testing.T) { + tests := map[string]string{ + "https://client.example/callback?source=cxdb": "https://client.example", + "http://localhost:6276/oauth/callback": "http://localhost:6276", + "http://127.0.0.1:6276/oauth/callback": "http://127.0.0.1:6276", + "http://[::1]:6276/oauth/callback": "http://[::1]:6276", + } + for redirectURI, origin := range tests { + t.Run(origin, func(t *testing.T) { + csp, err := oauthConsentCSP(redirectURI) + if err != nil { + t.Fatal(err) + } + want := "default-src 'none'; form-action 'self' " + origin + "; frame-ancestors 'none'; base-uri 'none'" + if csp != want { + t.Fatalf("consent CSP = %q, want %q", csp, want) + } + }) + } +} + +func TestOAuthAuthorizationRequiresState(t *testing.T) { + server, _ := testOAuthServer(t) + query := url.Values{ + "response_type": {"code"}, "client_id": {"client"}, + "redirect_uri": {"http://127.0.0.1:49152/callback"}, + "code_challenge": {"challenge"}, "code_challenge_method": {"S256"}, + } + if _, err := server.parseAuthorizationRequest(query); err == nil { + t.Fatal("authorization request without state was accepted") + } +} + +func TestOAuthRegistrationRejectsClientsAtPersistentCap(t *testing.T) { + server, _ := testOAuthServer(t) + for i := 0; i < maxOAuthClients; i++ { + if _, err := server.store.db.Exec(` + INSERT INTO oauth_clients (client_id, client_name, redirect_uris_json, created_at) + VALUES (?, ?, ?, ?) + `, fmt.Sprintf("cap-%d", i), "cap test", `["http://127.0.0.1/callback"]`, time.Now().UTC()); err != nil { + t.Fatalf("fill client cap at %d: %v", i, err) + } + } + request := httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(`{"redirect_uris":["http://127.0.0.1/callback"]}`)) + response := httptest.NewRecorder() + server.RegisterHandler(response, request) + if response.Code != http.StatusTooManyRequests || !strings.Contains(response.Body.String(), "registration_limit_reached") { + t.Fatalf("registration at cap status=%d body=%s", response.Code, response.Body.String()) + } +} + +func registerTestClient(t *testing.T, server *OAuthServer, redirect string) (string, string) { + t.Helper() + body, _ := json.Marshal(map[string]any{"client_name": "test", "redirect_uris": []string{redirect}, "token_endpoint_auth_method": "none"}) + request := httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(string(body))) + response := httptest.NewRecorder() + server.RegisterHandler(response, request) + if response.Code != http.StatusCreated { + t.Fatalf("register status = %d, body=%s", response.Code, response.Body.String()) + } + var payload struct { + ClientID string `json:"client_id"` + } + if err := json.Unmarshal(response.Body.Bytes(), &payload); err != nil { + t.Fatal(err) + } + return payload.ClientID, redirect +} + +func setSessionCookie(t *testing.T, store *SessionStore, sessionID string, request *http.Request) { + t.Helper() + recorder := httptest.NewRecorder() + store.SetCookie(recorder, sessionID) + request.AddCookie(recorder.Result().Cookies()[0]) +} + +func strconvQuote(value string) string { + raw, _ := json.Marshal(value) + return string(raw) +} diff --git a/gateway/pkg/auth/session.go b/gateway/pkg/auth/session.go index 2fa8569..c378b78 100644 --- a/gateway/pkg/auth/session.go +++ b/gateway/pkg/auth/session.go @@ -10,6 +10,7 @@ import ( "crypto/sha256" "database/sql" "encoding/hex" + "encoding/json" "errors" "fmt" "log" @@ -24,12 +25,33 @@ import ( // Session captures the authenticated user for a browser. type Session struct { - ID string - Email string - Name string - Picture string - CreatedAt time.Time - ExpiresAt time.Time + ID string + Email string + Name string + Picture string + Scopes []string + CreatedAt time.Time + ExpiresAt time.Time + AuthMethod string // Authentication method, for example "oidc" or "k8s_oidc". + Issuer string // Token issuer URL + Subject string // Stable subject within Issuer +} + +// HasScope returns true if the session includes the given scope. +func (s *Session) HasScope(scope string) bool { + for _, sc := range s.Scopes { + if sc == scope { + return true + } + } + return false +} + +// IsAPIToken reports whether this session came from a personal API token. +// Handlers can use this to prevent a bearer token from creating or revoking +// other personal credentials. +func (s *Session) IsAPIToken() bool { + return s != nil && s.AuthMethod == APITokenAuthMethod } // SessionStore handles persistence of sessions in SQLite and @@ -84,48 +106,105 @@ func (s *SessionStore) ensureSchema() error { expires_at TIMESTAMP NOT NULL ); CREATE INDEX IF NOT EXISTS idx_sessions_email ON sessions(email); + CREATE TABLE IF NOT EXISTS api_tokens ( + id TEXT PRIMARY KEY, + name TEXT NOT NULL, + issuer TEXT NOT NULL, + subject TEXT NOT NULL, + scopes TEXT NOT NULL, + token_hash TEXT NOT NULL UNIQUE, + created_at TIMESTAMP NOT NULL, + expires_at TIMESTAMP NOT NULL, + revoked_at TIMESTAMP, + last_used_at TIMESTAMP + ); + CREATE INDEX IF NOT EXISTS idx_api_tokens_owner ON api_tokens(issuer, subject); + CREATE INDEX IF NOT EXISTS idx_api_tokens_hash ON api_tokens(token_hash); ` if _, err := s.db.Exec(schema); err != nil { return fmt.Errorf("init schema: %w", err) } // Backfill for older schemas missing the picture column; ignore duplicate errors. _, _ = s.db.Exec(`ALTER TABLE sessions ADD COLUMN picture TEXT;`) + // These columns are deliberately nullable so that existing installations can + // be upgraded without rewriting or invalidating their browser sessions. + for _, statement := range []string{ + `ALTER TABLE sessions ADD COLUMN issuer TEXT`, + `ALTER TABLE sessions ADD COLUMN subject TEXT`, + `ALTER TABLE sessions ADD COLUMN scopes TEXT`, + `ALTER TABLE sessions ADD COLUMN auth_method TEXT`, + } { + _, _ = s.db.Exec(statement) + } return nil } // Create inserts a new session and returns its ID. func (s *SessionStore) Create(ctx context.Context, email, name, picture string) (string, error) { + return s.CreateForIdentity(ctx, "https://accounts.google.com", email, email, name, picture, "google_oauth", []string{"cxdb:read", "cxdb:write"}) +} + +// CreateForIdentity inserts a browser session with a stable issuer/subject +// identity and authorization scopes. Create remains the compatibility API for +// callers that only have a Google profile. +func (s *SessionStore) CreateForIdentity(ctx context.Context, issuer, subject, email, name, picture, authMethod string, scopes []string) (string, error) { id, err := randomID() if err != nil { return "", err } now := time.Now().UTC() expires := now.Add(s.ttl) + scopeJSON, err := json.Marshal(normalizeScopes(scopes)) + if err != nil { + return "", fmt.Errorf("encode session scopes: %w", err) + } _, err = s.db.ExecContext(ctx, ` - INSERT INTO sessions (id, email, name, picture, created_at, expires_at) - VALUES (?, ?, ?, ?, ?, ?) - `, id, email, name, picture, now, expires) + INSERT INTO sessions (id, email, name, picture, issuer, subject, scopes, auth_method, created_at, expires_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, id, email, name, picture, strings.TrimSpace(issuer), strings.TrimSpace(subject), string(scopeJSON), strings.TrimSpace(authMethod), now, expires) if err != nil { return "", fmt.Errorf("insert session: %w", err) } return id, nil } +// CreateWithIdentity is retained as a convenience for callers using the +// original field-oriented order introduced during the identity migration. +func (s *SessionStore) CreateWithIdentity(ctx context.Context, email, name, picture, issuer, subject string, scopes []string, authMethod string) (string, error) { + return s.CreateForIdentity(ctx, issuer, subject, email, name, picture, authMethod, scopes) +} + // Get returns a valid, non-expired session by ID. func (s *SessionStore) Get(ctx context.Context, id string) (*Session, error) { row := s.db.QueryRowContext(ctx, ` - SELECT id, email, name, picture, created_at, expires_at + SELECT id, email, name, picture, issuer, subject, scopes, auth_method, created_at, expires_at FROM sessions WHERE id = ? `, id) var sess Session - if err := row.Scan(&sess.ID, &sess.Email, &sess.Name, &sess.Picture, &sess.CreatedAt, &sess.ExpiresAt); err != nil { + var email, name, picture, issuer, subject, scopesJSON, authMethod sql.NullString + if err := row.Scan(&sess.ID, &email, &name, &picture, &issuer, &subject, &scopesJSON, &authMethod, &sess.CreatedAt, &sess.ExpiresAt); err != nil { if errors.Is(err, sql.ErrNoRows) { return nil, nil } return nil, fmt.Errorf("select session: %w", err) } + sess.Email, sess.Name, sess.Picture = email.String, name.String, picture.String + sess.Issuer, sess.Subject, sess.AuthMethod = issuer.String, subject.String, authMethod.String + if scopesJSON.Valid && scopesJSON.String != "" { + if err := json.Unmarshal([]byte(scopesJSON.String), &sess.Scopes); err != nil { + return nil, fmt.Errorf("decode session scopes: %w", err) + } + } + // Sessions that predate the identity migration were authenticated Google + // browser sessions. Preserve them until their normal expiry. + if !issuer.Valid && !subject.Valid && !scopesJSON.Valid && !authMethod.Valid { + sess.Issuer = "https://accounts.google.com" + sess.Subject = sess.Email + sess.AuthMethod = "google_oauth" + sess.Scopes = []string{"cxdb:read", "cxdb:write"} + } if time.Now().After(sess.ExpiresAt) { _ = s.Delete(ctx, id) return nil, nil @@ -194,6 +273,22 @@ func (s *SessionStore) SetCookie(w http.ResponseWriter, sessionID string) { }) } +// CSRFToken returns a session-bound token for browser credential-management requests. +func (s *SessionStore) CSRFToken(session *Session) string { + if session == nil || session.ID == "" { + return "" + } + mac := hmac.New(sha256.New, s.secret) + _, _ = mac.Write([]byte("cxdb-csrf\x00" + session.ID)) + return hex.EncodeToString(mac.Sum(nil)) +} + +// ValidCSRFToken checks a session-bound CSRF token in constant time. +func (s *SessionStore) ValidCSRFToken(session *Session, token string) bool { + expected := s.CSRFToken(session) + return expected != "" && hmac.Equal([]byte(expected), []byte(strings.TrimSpace(token))) +} + // ClearCookie removes the session cookie from the browser. func (s *SessionStore) ClearCookie(w http.ResponseWriter) { http.SetCookie(w, &http.Cookie{ diff --git a/gateway/pkg/auth/session_test.go b/gateway/pkg/auth/session_test.go new file mode 100644 index 0000000..8d07159 --- /dev/null +++ b/gateway/pkg/auth/session_test.go @@ -0,0 +1,38 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + +package auth + +import ( + "context" + "path/filepath" + "testing" + "time" +) + +func TestLegacyBrowserSessionKeepsDefaultIdentityAndScopes(t *testing.T) { + store, err := NewSessionStore(filepath.Join(t.TempDir(), "sessions.sqlite"), "session", time.Hour, "", false, "test-secret") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = store.Close() }) + now := time.Now().UTC() + _, err = store.db.ExecContext(context.Background(), ` + INSERT INTO sessions (id, email, name, picture, created_at, expires_at) + VALUES (?, ?, ?, ?, ?, ?) + `, "legacy", "legacy@example.com", "Legacy", "", now, now.Add(time.Hour)) + if err != nil { + t.Fatal(err) + } + + session, err := store.Get(context.Background(), "legacy") + if err != nil { + t.Fatal(err) + } + if session == nil || session.Issuer != "https://accounts.google.com" || session.Subject != "legacy@example.com" || session.AuthMethod != "google_oauth" { + t.Fatalf("legacy identity = %+v", session) + } + if !session.HasScope("cxdb:read") || !session.HasScope("cxdb:write") { + t.Fatalf("legacy scopes = %v", session.Scopes) + } +} diff --git a/gateway/pkg/mcpserver/server.go b/gateway/pkg/mcpserver/server.go new file mode 100644 index 0000000..bab8a15 --- /dev/null +++ b/gateway/pkg/mcpserver/server.go @@ -0,0 +1,315 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + +// Package mcpserver exposes CXDB operations through remote Streamable HTTP MCP. +package mcpserver + +import ( + "bytes" + "context" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "io" + "log/slog" + "net/http" + "net/url" + "strconv" + "strings" + "time" + + mcpauth "github.com/modelcontextprotocol/go-sdk/auth" + "github.com/modelcontextprotocol/go-sdk/mcp" + cxdbauth "github.com/strongdm/cxdb/gateway/pkg/auth" + "github.com/vmihailenco/msgpack/v5" +) + +const maxBackendResponse = 8 << 20 + +// New returns a bearer-protected, origin-protected MCP Streamable HTTP handler. +func New(backendURL, resourceMetadataURL string, verifiers []cxdbauth.BearerTokenVerifier, logger *slog.Logger) (http.Handler, error) { + backend, err := url.Parse(strings.TrimSuffix(backendURL, "/")) + if err != nil || backend.Scheme == "" || backend.Host == "" { + return nil, errors.New("invalid CXDB backend URL") + } + api := &backendClient{base: backend, client: &http.Client{Timeout: 30 * time.Second}} + server := mcp.NewServer(&mcp.Implementation{Name: "cxdb", Version: "0.1.0"}, &mcp.ServerOptions{ + Instructions: "Read and append CXDB Turn DAG contexts. Read exact turns before treating bounded summaries as complete.", + Logger: logger, + }) + registerTools(server, api) + + stream := mcp.NewStreamableHTTPHandler(func(*http.Request) *mcp.Server { return server }, &mcp.StreamableHTTPOptions{ + Stateless: true, + PropagateRequestCancellation: true, + MaxRequestBodyBytes: 1 << 20, + Logger: logger, + }) + originProtected := http.NewCrossOriginProtection().Handler(stream) + verifier := func(ctx context.Context, token string, req *http.Request) (*mcpauth.TokenInfo, error) { + for _, candidate := range verifiers { + var session *cxdbauth.Session + var verifyErr error + if requestVerifier, ok := candidate.(cxdbauth.RequestTokenVerifier); ok { + session, verifyErr = requestVerifier.VerifyWithRequest(req, token) + } else { + session, verifyErr = candidate.Verify(token) + } + if verifyErr == nil && session != nil { + return &mcpauth.TokenInfo{ + Scopes: session.Scopes, Expiration: session.ExpiresAt, UserID: session.Issuer + "|" + session.Subject, + }, nil + } + } + return nil, fmt.Errorf("%w: bearer token is invalid", mcpauth.ErrInvalidToken) + } + return mcpauth.RequireBearerToken(verifier, &mcpauth.RequireBearerTokenOptions{ + ResourceMetadataURL: resourceMetadataURL, + })(originProtected), nil +} + +type backendClient struct { + base *url.URL + client *http.Client +} + +type listInput struct { + Limit int `json:"limit,omitempty"` +} + +type searchInput struct { + Query string `json:"query"` + Limit int `json:"limit,omitempty"` +} + +type contextInput struct { + ContextID string `json:"context_id"` +} + +type turnsInput struct { + ContextID string `json:"context_id"` + Limit int `json:"limit,omitempty"` + BeforeTurn string `json:"before_turn_id,omitempty"` + ExactTurnID string `json:"turn_id,omitempty"` +} + +type createInput struct { + BaseTurnID string `json:"base_turn_id,omitempty"` +} + +type appendMessageInput struct { + ContextID string `json:"context_id"` + Role string `json:"role"` + Text string `json:"text"` +} + +type appendRawInput struct { + ContextID string `json:"context_id"` + TypeID string `json:"type_id"` + TypeVersion uint32 `json:"type_version"` + PayloadBase64 string `json:"payload_base64"` +} + +func registerTools(server *mcp.Server, api *backendClient) { + mcp.AddTool(server, &mcp.Tool{Name: "cxdb_list_contexts", Description: "List recent CXDB contexts."}, func(ctx context.Context, _ *mcp.CallToolRequest, input listInput) (*mcp.CallToolResult, map[string]any, error) { + if err := requireScope(ctx, "cxdb:read"); err != nil { + return nil, nil, err + } + limit := boundedLimit(input.Limit) + return api.call(ctx, http.MethodGet, "/v1/contexts?limit="+strconv.Itoa(limit), nil) + }) + mcp.AddTool(server, &mcp.Tool{Name: "cxdb_search_contexts", Description: "Search contexts with CXDB Query Language."}, func(ctx context.Context, _ *mcp.CallToolRequest, input searchInput) (*mcp.CallToolResult, map[string]any, error) { + if err := requireScope(ctx, "cxdb:read"); err != nil { + return nil, nil, err + } + if strings.TrimSpace(input.Query) == "" { + return nil, nil, errors.New("query is required") + } + path := "/v1/contexts/search?q=" + url.QueryEscape(input.Query) + "&limit=" + strconv.Itoa(boundedLimit(input.Limit)) + return api.call(ctx, http.MethodGet, path, nil) + }) + mcp.AddTool(server, &mcp.Tool{Name: "cxdb_get_context", Description: "Get one context head and metadata."}, func(ctx context.Context, _ *mcp.CallToolRequest, input contextInput) (*mcp.CallToolResult, map[string]any, error) { + if err := requireScope(ctx, "cxdb:read"); err != nil { + return nil, nil, err + } + if err := numericID(input.ContextID); err != nil { + return nil, nil, err + } + return api.call(ctx, http.MethodGet, "/v1/contexts/"+input.ContextID, nil) + }) + mcp.AddTool(server, &mcp.Tool{Name: "cxdb_get_turns", Description: "Read typed turns. Set turn_id to hydrate one exact complete turn."}, func(ctx context.Context, _ *mcp.CallToolRequest, input turnsInput) (*mcp.CallToolResult, map[string]any, error) { + if err := requireScope(ctx, "cxdb:read"); err != nil { + return nil, nil, err + } + if err := numericID(input.ContextID); err != nil { + return nil, nil, err + } + query := url.Values{"limit": {strconv.Itoa(boundedLimit(input.Limit))}, "view": {"typed"}} + if input.BeforeTurn != "" { + if err := numericID(input.BeforeTurn); err != nil { + return nil, nil, err + } + query.Set("before_turn_id", input.BeforeTurn) + } + if input.ExactTurnID != "" { + if err := numericID(input.ExactTurnID); err != nil { + return nil, nil, err + } + query.Set("turn_id", input.ExactTurnID) + } + return api.call(ctx, http.MethodGet, "/v1/contexts/"+input.ContextID+"/turns?"+query.Encode(), nil) + }) + mcp.AddTool(server, &mcp.Tool{Name: "cxdb_get_provenance", Description: "Get provenance for one context."}, func(ctx context.Context, _ *mcp.CallToolRequest, input contextInput) (*mcp.CallToolResult, map[string]any, error) { + if err := requireScope(ctx, "cxdb:read"); err != nil { + return nil, nil, err + } + if err := numericID(input.ContextID); err != nil { + return nil, nil, err + } + return api.call(ctx, http.MethodGet, "/v1/contexts/"+input.ContextID+"/provenance", nil) + }) + mcp.AddTool(server, &mcp.Tool{Name: "cxdb_create_context", Description: "Create a new context or fork from a turn."}, func(ctx context.Context, _ *mcp.CallToolRequest, input createInput) (*mcp.CallToolResult, map[string]any, error) { + if err := requireScope(ctx, "cxdb:write"); err != nil { + return nil, nil, err + } + body := map[string]any{} + if input.BaseTurnID != "" { + if err := numericID(input.BaseTurnID); err != nil { + return nil, nil, err + } + body["base_turn_id"] = input.BaseTurnID + } + return api.call(ctx, http.MethodPost, "/v1/contexts/create", body) + }) + mcp.AddTool(server, &mcp.Tool{Name: "cxdb_append_message", Description: "Append a canonical user, assistant, or system message."}, func(ctx context.Context, _ *mcp.CallToolRequest, input appendMessageInput) (*mcp.CallToolResult, map[string]any, error) { + if err := requireScope(ctx, "cxdb:write"); err != nil { + return nil, nil, err + } + if err := numericID(input.ContextID); err != nil { + return nil, nil, err + } + payload, err := canonicalMessage(input.Role, input.Text) + if err != nil { + return nil, nil, err + } + body := map[string]any{"type_id": "cxdb.ConversationItem", "type_version": 3, "payload_base64": base64.StdEncoding.EncodeToString(payload)} + return api.call(ctx, http.MethodPost, "/v1/contexts/"+input.ContextID+"/append", body) + }) + mcp.AddTool(server, &mcp.Tool{Name: "cxdb_append_turn", Description: "Append a raw MessagePack turn with an explicit registered type."}, func(ctx context.Context, _ *mcp.CallToolRequest, input appendRawInput) (*mcp.CallToolResult, map[string]any, error) { + if err := requireScope(ctx, "cxdb:write"); err != nil { + return nil, nil, err + } + if err := numericID(input.ContextID); err != nil { + return nil, nil, err + } + if input.TypeID == "" || input.TypeVersion == 0 { + return nil, nil, errors.New("type_id and type_version are required") + } + decoded, err := base64.StdEncoding.DecodeString(input.PayloadBase64) + if err != nil || len(decoded) > 4<<20 { + return nil, nil, errors.New("payload_base64 must contain at most 4 MiB") + } + body := map[string]any{"type_id": input.TypeID, "type_version": input.TypeVersion, "payload_base64": input.PayloadBase64} + return api.call(ctx, http.MethodPost, "/v1/contexts/"+input.ContextID+"/append", body) + }) +} + +func (c *backendClient) call(ctx context.Context, method, path string, body any) (*mcp.CallToolResult, map[string]any, error) { + requestURL := *c.base + parsed, err := url.Parse(path) + if err != nil { + return nil, nil, err + } + requestURL.Path = parsed.Path + requestURL.RawQuery = parsed.RawQuery + var reader io.Reader + if body != nil { + raw, marshalErr := json.Marshal(body) + if marshalErr != nil { + return nil, nil, marshalErr + } + reader = bytes.NewReader(raw) + } + req, err := http.NewRequestWithContext(ctx, method, requestURL.String(), reader) + if err != nil { + return nil, nil, err + } + if body != nil { + req.Header.Set("Content-Type", "application/json") + } + resp, err := c.client.Do(req) + if err != nil { + return nil, nil, fmt.Errorf("CXDB backend request: %w", err) + } + defer func() { _ = resp.Body.Close() }() + raw, err := io.ReadAll(io.LimitReader(resp.Body, maxBackendResponse+1)) + if err != nil { + return nil, nil, err + } + if len(raw) > maxBackendResponse { + return nil, nil, errors.New("CXDB backend response exceeds 8 MiB") + } + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + return nil, nil, fmt.Errorf("CXDB backend returned %d: %s", resp.StatusCode, strings.TrimSpace(string(raw))) + } + var decoded any + if err := json.Unmarshal(raw, &decoded); err != nil { + return nil, nil, fmt.Errorf("decode CXDB response: %w", err) + } + return nil, map[string]any{"response": decoded}, nil +} + +func requireScope(ctx context.Context, scope string) error { + info := mcpauth.TokenInfoFromContext(ctx) + if info == nil { + return errors.New("authentication is required") + } + for _, granted := range info.Scopes { + if granted == scope { + return nil + } + } + return fmt.Errorf("insufficient scope: %s is required", scope) +} + +func numericID(value string) error { + if value == "" { + return errors.New("ID is required") + } + if _, err := strconv.ParseUint(value, 10, 64); err != nil { + return errors.New("ID must be an unsigned integer") + } + return nil +} + +func boundedLimit(value int) int { + if value <= 0 { + return 50 + } + if value > 200 { + return 200 + } + return value +} + +func canonicalMessage(role, text string) ([]byte, error) { + if text == "" { + return nil, errors.New("text is required") + } + item := map[int]any{2: "complete", 3: time.Now().UnixMilli()} + switch role { + case "user": + item[1] = "user_input" + item[10] = map[int]any{1: text} + case "assistant": + item[1] = "assistant_turn" + item[11] = map[int]any{1: text} + case "system": + item[1] = "system" + item[12] = map[int]any{1: "info", 3: text} + default: + return nil, errors.New("role must be user, assistant, or system") + } + return msgpack.Marshal(item) +} diff --git a/gateway/pkg/mcpserver/server_test.go b/gateway/pkg/mcpserver/server_test.go new file mode 100644 index 0000000..71d768f --- /dev/null +++ b/gateway/pkg/mcpserver/server_test.go @@ -0,0 +1,296 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + +package mcpserver + +import ( + "context" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "html" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "net/url" + "path/filepath" + "regexp" + "strings" + "testing" + "time" + + "github.com/modelcontextprotocol/go-sdk/mcp" + cxdbauth "github.com/strongdm/cxdb/gateway/pkg/auth" + "github.com/vmihailenco/msgpack/v5" +) + +type staticVerifier struct{ scopes []string } + +func (v staticVerifier) Verify(token string) (*cxdbauth.Session, error) { + if token != "test-token" { + return nil, cxdbauth.ErrAPITokenNotFound + } + return &cxdbauth.Session{ID: "test", Issuer: "test", Subject: "user", Email: "user@example.com", Scopes: v.scopes, ExpiresAt: time.Now().Add(time.Hour)}, nil +} + +type bearerTransport struct { + base http.RoundTripper + token string +} + +func (t bearerTransport) RoundTrip(request *http.Request) (*http.Response, error) { + clone := request.Clone(request.Context()) + clone.Header = request.Header.Clone() + clone.Header.Set("Authorization", "Bearer "+t.token) + return t.base.RoundTrip(clone) +} + +func TestOfficialClientHandshakeReadAndWriteTools(t *testing.T) { + var appended bool + backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + switch { + case r.Method == http.MethodGet && r.URL.Path == "/v1/contexts": + _, _ = io.WriteString(w, `{"contexts":[{"context_id":"1"}]}`) + case r.Method == http.MethodPost && r.URL.Path == "/v1/contexts/1/append": + var body map[string]any + if err := json.NewDecoder(r.Body).Decode(&body); err != nil { + t.Fatal(err) + } + if body["type_id"] != "cxdb.ConversationItem" || body["payload_base64"] == "" { + t.Fatalf("unexpected append body: %#v", body) + } + appended = true + _, _ = io.WriteString(w, `{"turn_id":"2"}`) + default: + http.NotFound(w, r) + } + })) + t.Cleanup(backend.Close) + + handler, err := New(backend.URL, "https://cxdb.example/.well-known/oauth-protected-resource/mcp", []cxdbauth.BearerTokenVerifier{staticVerifier{scopes: []string{"cxdb:read", "cxdb:write"}}}, slog.Default()) + if err != nil { + t.Fatal(err) + } + remote := httptest.NewServer(handler) + t.Cleanup(remote.Close) + + client := mcp.NewClient(&mcp.Implementation{Name: "cxdb-test", Version: "1"}, nil) + httpClient := &http.Client{Transport: bearerTransport{base: http.DefaultTransport, token: "test-token"}} + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + session, err := client.Connect(ctx, &mcp.StreamableClientTransport{Endpoint: remote.URL, HTTPClient: httpClient, DisableStandaloneSSE: true}, nil) + if err != nil { + t.Fatalf("official MCP client handshake: %v", err) + } + defer func() { _ = session.Close() }() + if _, err := session.CallTool(ctx, &mcp.CallToolParams{Name: "cxdb_list_contexts", Arguments: map[string]any{"limit": 1}}); err != nil { + t.Fatalf("read tool: %v", err) + } + result, err := session.CallTool(ctx, &mcp.CallToolParams{Name: "cxdb_append_message", Arguments: map[string]any{"context_id": "1", "role": "user", "text": "hello"}}) + if err != nil { + t.Fatalf("write tool: %v", err) + } + if result.IsError || !appended { + t.Fatalf("write tool did not append: result=%+v appended=%v", result, appended) + } +} + +func TestWriteToolRequiresWriteScope(t *testing.T) { + backend := httptest.NewServer(http.NotFoundHandler()) + t.Cleanup(backend.Close) + handler, err := New(backend.URL, "https://cxdb.example/metadata", []cxdbauth.BearerTokenVerifier{staticVerifier{scopes: []string{"cxdb:read"}}}, slog.Default()) + if err != nil { + t.Fatal(err) + } + remote := httptest.NewServer(handler) + t.Cleanup(remote.Close) + client := mcp.NewClient(&mcp.Implementation{Name: "cxdb-test", Version: "1"}, nil) + session, err := client.Connect(context.Background(), &mcp.StreamableClientTransport{Endpoint: remote.URL, HTTPClient: &http.Client{Transport: bearerTransport{base: http.DefaultTransport, token: "test-token"}}, DisableStandaloneSSE: true}, nil) + if err != nil { + t.Fatal(err) + } + defer func() { _ = session.Close() }() + result, err := session.CallTool(context.Background(), &mcp.CallToolParams{Name: "cxdb_create_context", Arguments: map[string]any{}}) + if err != nil { + t.Fatal(err) + } + if !result.IsError { + t.Fatal("read-only token was allowed to call a write tool") + } +} + +func TestWriteOnlyTokenCanConnectAndUseWriteTool(t *testing.T) { + var created bool + backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodPost && r.URL.Path == "/v1/contexts/create" { + created = true + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"context_id":"1"}`) + return + } + http.NotFound(w, r) + })) + t.Cleanup(backend.Close) + handler, err := New(backend.URL, "https://cxdb.example/metadata", []cxdbauth.BearerTokenVerifier{staticVerifier{scopes: []string{"cxdb:write"}}}, slog.Default()) + if err != nil { + t.Fatal(err) + } + remote := httptest.NewServer(handler) + t.Cleanup(remote.Close) + client := mcp.NewClient(&mcp.Implementation{Name: "cxdb-test", Version: "1"}, nil) + session, err := client.Connect(context.Background(), &mcp.StreamableClientTransport{Endpoint: remote.URL, HTTPClient: &http.Client{Transport: bearerTransport{base: http.DefaultTransport, token: "test-token"}}, DisableStandaloneSSE: true}, nil) + if err != nil { + t.Fatal(err) + } + defer func() { _ = session.Close() }() + result, err := session.CallTool(context.Background(), &mcp.CallToolParams{Name: "cxdb_create_context", Arguments: map[string]any{}}) + if err != nil { + t.Fatal(err) + } + if result.IsError || !created { + t.Fatalf("write-only token did not create context: result=%+v created=%v", result, created) + } +} + +func TestCanonicalMessagesUseNumericTags(t *testing.T) { + tests := []struct { + role string + variantTag int + textTag int + }{ + {role: "user", variantTag: 10, textTag: 1}, + {role: "assistant", variantTag: 11, textTag: 1}, + {role: "system", variantTag: 12, textTag: 3}, + } + for _, test := range tests { + t.Run(test.role, func(t *testing.T) { + payload, err := canonicalMessage(test.role, "hello") + if err != nil { + t.Fatal(err) + } + var item map[int]msgpack.RawMessage + if err := msgpack.Unmarshal(payload, &item); err != nil { + t.Fatal(err) + } + var variant map[int]string + if err := msgpack.Unmarshal(item[test.variantTag], &variant); err != nil { + t.Fatalf("decode variant tag %d: %v", test.variantTag, err) + } + if got := variant[test.textTag]; got != "hello" { + t.Fatalf("text tag %d = %#v", test.textTag, got) + } + if _, exists := item[0]; exists { + t.Fatal("unexpected zero tag") + } + }) + } +} + +func TestOAuthAccessTokenConnectsWithOfficialClient(t *testing.T) { + store, err := cxdbauth.NewSessionStore(filepath.Join(t.TempDir(), "oauth.sqlite"), "session", time.Hour, "", false, "integration-secret") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = store.Close() }) + oauth, err := cxdbauth.NewOAuthServer(store, "https://cxdb.example", "/auth/login") + if err != nil { + t.Fatal(err) + } + + redirectURI := "http://127.0.0.1:49152/callback" + registrationBody := `{"client_name":"official MCP client","redirect_uris":["` + redirectURI + `"],"token_endpoint_auth_method":"none"}` + registrationRequest := httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(registrationBody)) + registrationResponse := httptest.NewRecorder() + oauth.RegisterHandler(registrationResponse, registrationRequest) + if registrationResponse.Code != http.StatusCreated { + t.Fatalf("register status = %d, body=%s", registrationResponse.Code, registrationResponse.Body.String()) + } + var registration struct { + ClientID string `json:"client_id"` + } + if err := json.Unmarshal(registrationResponse.Body.Bytes(), ®istration); err != nil { + t.Fatal(err) + } + + sessionID, err := store.CreateForIdentity(t.Context(), "https://id.example", "alice", "alice@example.com", "Alice", "", "oidc", []string{"cxdb:read", "cxdb:write"}) + if err != nil { + t.Fatal(err) + } + verifier := strings.Repeat("v", 64) + digest := sha256.Sum256([]byte(verifier)) + query := url.Values{ + "response_type": {"code"}, + "client_id": {registration.ClientID}, + "redirect_uri": {redirectURI}, + "state": {"client-state"}, + "scope": {"cxdb:read cxdb:write"}, + "resource": {"https://cxdb.example/mcp"}, + "code_challenge": {base64.RawURLEncoding.EncodeToString(digest[:])}, + "code_challenge_method": {"S256"}, + } + authorizeRequest := httptest.NewRequest(http.MethodGet, "/oauth/authorize?"+query.Encode(), nil) + addSessionCookie(store, sessionID, authorizeRequest) + authorizeResponse := httptest.NewRecorder() + oauth.AuthorizeHandler(authorizeResponse, authorizeRequest) + match := regexp.MustCompile(`name="request" value="([^"]+)"`).FindStringSubmatch(authorizeResponse.Body.String()) + if authorizeResponse.Code != http.StatusOK || len(match) != 2 { + t.Fatalf("authorization consent status = %d, body=%s", authorizeResponse.Code, authorizeResponse.Body.String()) + } + consentForm := url.Values{"request": {html.UnescapeString(match[1])}, "decision": {"allow"}} + consentRequest := httptest.NewRequest(http.MethodPost, "/oauth/authorize", strings.NewReader(consentForm.Encode())) + consentRequest.Header.Set("Content-Type", "application/x-www-form-urlencoded") + addSessionCookie(store, sessionID, consentRequest) + consentResponse := httptest.NewRecorder() + oauth.AuthorizeHandler(consentResponse, consentRequest) + location, err := url.Parse(consentResponse.Header().Get("Location")) + if err != nil || location.Query().Get("code") == "" { + t.Fatalf("consent redirect = %q, err=%v", consentResponse.Header().Get("Location"), err) + } + tokenForm := url.Values{ + "grant_type": {"authorization_code"}, + "code": {location.Query().Get("code")}, + "client_id": {registration.ClientID}, + "redirect_uri": {redirectURI}, + "code_verifier": {verifier}, + } + tokenRequest := httptest.NewRequest(http.MethodPost, "/oauth/token", strings.NewReader(tokenForm.Encode())) + tokenRequest.Header.Set("Content-Type", "application/x-www-form-urlencoded") + tokenResponse := httptest.NewRecorder() + oauth.TokenHandler(tokenResponse, tokenRequest) + var tokenPayload struct { + AccessToken string `json:"access_token"` + } + if tokenResponse.Code != http.StatusOK || json.Unmarshal(tokenResponse.Body.Bytes(), &tokenPayload) != nil || tokenPayload.AccessToken == "" { + t.Fatalf("token response status = %d, body=%s", tokenResponse.Code, tokenResponse.Body.String()) + } + + backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"contexts":[]}`) + })) + t.Cleanup(backend.Close) + handler, err := New(backend.URL, "https://cxdb.example/.well-known/oauth-protected-resource/mcp", []cxdbauth.BearerTokenVerifier{oauth}, slog.Default()) + if err != nil { + t.Fatal(err) + } + remote := httptest.NewServer(handler) + t.Cleanup(remote.Close) + client := mcp.NewClient(&mcp.Implementation{Name: "cxdb-oauth-test", Version: "1"}, nil) + httpClient := &http.Client{Transport: bearerTransport{base: http.DefaultTransport, token: tokenPayload.AccessToken}} + mcpSession, err := client.Connect(t.Context(), &mcp.StreamableClientTransport{Endpoint: remote.URL, HTTPClient: httpClient, DisableStandaloneSSE: true}, nil) + if err != nil { + t.Fatalf("OAuth-backed official MCP client handshake: %v", err) + } + defer func() { _ = mcpSession.Close() }() + if _, err := mcpSession.CallTool(t.Context(), &mcp.CallToolParams{Name: "cxdb_list_contexts", Arguments: map[string]any{"limit": 1}}); err != nil { + t.Fatalf("OAuth-backed read tool: %v", err) + } +} + +func addSessionCookie(store *cxdbauth.SessionStore, sessionID string, request *http.Request) { + recorder := httptest.NewRecorder() + store.SetCookie(recorder, sessionID) + request.AddCookie(recorder.Result().Cookies()[0]) +} diff --git a/gateway/pkg/proxy/peer.go b/gateway/pkg/proxy/peer.go new file mode 100644 index 0000000..a731e22 --- /dev/null +++ b/gateway/pkg/proxy/peer.go @@ -0,0 +1,62 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + +package proxy + +import ( + "net" + "net/http" + "strings" +) + +// observedPeerIP returns the TCP peer's address parsed from req.RemoteAddr. +// +// This helper is the sole source of "real client address" at the gateway +// trust boundary. It NEVER reads any HTTP header. Every gateway site that +// needs a real peer (reverse-proxy XFF write, request logging, rate-limit +// bucket key, debug-auth IP allowlist) MUST call this helper and MUST NOT +// inspect `X-Forwarded-For` or `Forwarded` directly — those headers are +// attacker-controllable (see ADR-006). +// +// If `RemoteAddr` lacks a port (unusual — net/http populates it with +// "host:port"), the raw value is returned unchanged rather than mangled. +func observedPeerIP(req *http.Request) string { + host, _, err := net.SplitHostPort(req.RemoteAddr) + if err != nil { + return req.RemoteAddr + } + return host +} + +// rateLimitClientIP uses X-Forwarded-For only when the TCP peer is in an +// explicit trusted-proxy network. It walks the chain from right to left and +// returns the first untrusted address, which prevents client-supplied prefixes +// from changing the bucket key. +func rateLimitClientIP(req *http.Request, trusted []*net.IPNet) string { + peer := observedPeerIP(req) + peerIP := net.ParseIP(peer) + if peerIP == nil || !ipInNetworks(peerIP, trusted) { + return peer + } + forwarded := strings.Join(req.Header.Values("X-Forwarded-For"), ",") + parts := strings.Split(forwarded, ",") + for index := len(parts) - 1; index >= 0; index-- { + candidate := net.ParseIP(strings.TrimSpace(parts[index])) + if candidate == nil { + return peer + } + if !ipInNetworks(candidate, trusted) { + return candidate.String() + } + } + return peer +} + +func ipInNetworks(ip net.IP, networks []*net.IPNet) bool { + for _, network := range networks { + if network.Contains(ip) { + return true + } + } + return false +} diff --git a/gateway/pkg/proxy/peer_test.go b/gateway/pkg/proxy/peer_test.go new file mode 100644 index 0000000..8e2ce62 --- /dev/null +++ b/gateway/pkg/proxy/peer_test.go @@ -0,0 +1,38 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + +package proxy + +import ( + "net" + "net/http/httptest" + "testing" +) + +func TestRateLimitClientIPTrustBoundary(t *testing.T) { + _, trusted, err := net.ParseCIDR("10.0.0.0/8") + if err != nil { + t.Fatal(err) + } + tests := []struct { + name string + remoteAddr string + forwarded string + want string + }{ + {name: "direct client ignores header", remoteAddr: "198.51.100.7:1234", forwarded: "203.0.113.9", want: "198.51.100.7"}, + {name: "trusted proxy uses client", remoteAddr: "10.0.0.4:1234", forwarded: "203.0.113.9", want: "203.0.113.9"}, + {name: "trusted chain skips trusted hops", remoteAddr: "10.0.0.4:1234", forwarded: "192.0.2.8, 203.0.113.9, 10.1.2.3", want: "203.0.113.9"}, + {name: "malformed chain fails closed", remoteAddr: "10.0.0.4:1234", forwarded: "203.0.113.9, invalid", want: "10.0.0.4"}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + request := httptest.NewRequest("GET", "https://cxdb.example/login", nil) + request.RemoteAddr = test.remoteAddr + request.Header.Set("X-Forwarded-For", test.forwarded) + if got := rateLimitClientIP(request, []*net.IPNet{trusted}); got != test.want { + t.Fatalf("rateLimitClientIP() = %q, want %q", got, test.want) + } + }) + } +} diff --git a/gateway/pkg/proxy/reverse.go b/gateway/pkg/proxy/reverse.go index 3ffd7dc..1888883 100644 --- a/gateway/pkg/proxy/reverse.go +++ b/gateway/pkg/proxy/reverse.go @@ -27,37 +27,48 @@ func NewReverseProxy(backendURL string, logger *slog.Logger) (*ReverseProxy, err return nil, err } - proxy := httputil.NewSingleHostReverseProxy(target) - - // Custom director to set headers - originalDirector := proxy.Director - proxy.Director = func(req *http.Request) { - originalDirector(req) - - // Set the host to the target - req.Host = target.Host - - // Forward client IP - clientIP := extractClientIP(req) - if existing := req.Header.Get("X-Forwarded-For"); existing != "" { - req.Header.Set("X-Forwarded-For", existing+", "+clientIP) - } else { - req.Header.Set("X-Forwarded-For", clientIP) - } - - // Forward the original protocol - if req.Header.Get("X-Forwarded-Proto") == "" { - if req.TLS != nil { - req.Header.Set("X-Forwarded-Proto", "https") - } else { - req.Header.Set("X-Forwarded-Proto", "http") + proxy := &httputil.ReverseProxy{} + + // Use Rewrite (Go 1.20+). Rewrite is mutually exclusive with Director + // and disables the stdlib's default `X-Forwarded-For` auto-append — + // essential for the Sprint 019 / ADR-006 trust contract: the gateway + // MUST be the sole writer of `X-Forwarded-For` on outbound requests. + proxy.Rewrite = func(r *httputil.ProxyRequest) { + // Point the outbound request at the target backend. + r.SetURL(target) + r.Out.Host = target.Host + + // XFF trust contract (Sprint 019 / ADR-006): the gateway is the + // trust boundary. DROP any caller-supplied `X-Forwarded-For` and + // `Forwarded` headers first — they are attacker-controllable. + // Then set `X-Forwarded-For` to our own TCP-peer observation. The + // `observedPeerIP` helper is the single source of real-client-IP + // truth across the gateway (logging, rate-limit, this director). + r.Out.Header.Del("X-Forwarded-For") + r.Out.Header.Del("Forwarded") + // Identity headers are gateway assertions. Never forward caller values. + for header := range r.Out.Header { + if strings.HasPrefix(strings.ToLower(header), "x-cxdb-") { + r.Out.Header.Del(header) } } - - // Forward the original host - if req.Header.Get("X-Forwarded-Host") == "" { - req.Header.Set("X-Forwarded-Host", req.Host) + // Authentication is complete at the gateway. Do not forward browser + // cookies or bearer credentials to the Rust backend. + r.Out.Header.Del("Authorization") + r.Out.Header.Del("Cookie") + r.Out.Header.Set("X-Forwarded-For", observedPeerIP(r.In)) + + // X-Forwarded-Proto is a gateway assertion. Never preserve a caller value. + r.Out.Header.Del("X-Forwarded-Proto") + if r.In.TLS != nil { + r.Out.Header.Set("X-Forwarded-Proto", "https") + } else { + r.Out.Header.Set("X-Forwarded-Proto", "http") } + + // The Rust backend does not need the public host. Drop this header so a + // caller-selected Host value cannot cross the gateway trust boundary. + r.Out.Header.Del("X-Forwarded-Host") } // Custom error handler @@ -94,18 +105,3 @@ func (rp *ReverseProxy) ServeHTTP(w http.ResponseWriter, r *http.Request) { func (rp *ReverseProxy) Target() string { return rp.target.String() } - -func extractClientIP(r *http.Request) string { - // Check X-Forwarded-For first (in case we're behind another proxy) - if xff := r.Header.Get("X-Forwarded-For"); xff != "" { - parts := strings.Split(xff, ",") - return strings.TrimSpace(parts[0]) - } - - // Fall back to RemoteAddr - host, _, err := net.SplitHostPort(r.RemoteAddr) - if err != nil { - return r.RemoteAddr - } - return host -} diff --git a/gateway/pkg/proxy/reverse_security_test.go b/gateway/pkg/proxy/reverse_security_test.go new file mode 100644 index 0000000..194d7be --- /dev/null +++ b/gateway/pkg/proxy/reverse_security_test.go @@ -0,0 +1,64 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + +package proxy + +import ( + "log/slog" + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +func TestReverseProxyReplacesForwardedAndStripsCredentials(t *testing.T) { + seen := make(chan http.Header, 1) + backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + seen <- r.Header.Clone() + w.WriteHeader(http.StatusNoContent) + })) + defer backend.Close() + + reverse, err := NewReverseProxy(backend.URL, slog.Default()) + if err != nil { + t.Fatal(err) + } + req := httptest.NewRequest(http.MethodPost, "/v1/contexts/1/append", strings.NewReader("{}")) + req.RemoteAddr = "198.51.100.7:44321" + req.Host = "cxdb.example" + req.Header.Set("X-Forwarded-For", "203.0.113.9") + req.Header.Set("Forwarded", "for=203.0.113.9") + req.Header.Set("X-Forwarded-Proto", "https") + req.Header.Set("X-Forwarded-Host", "evil.example") + req.Header.Set("X-Cxdb-Writer-Method", "admin") + req.Header.Set("X-Cxdb-Writer-Subject", "admin") + req.Header.Set("X-Cxdb-Writer-Issuer", "admin") + req.Header.Set("X-Cxdb-User-Email", "admin@example.com") + req.Header.Set("X-Cxdb-Unknown", "admin") + req.Header.Set("Authorization", "Bearer gateway-token") + req.Header.Set("Cookie", "cxdb_session=browser-cookie") + + response := httptest.NewRecorder() + reverse.ServeHTTP(response, req) + if response.Code != http.StatusNoContent { + t.Fatalf("proxy status = %d", response.Code) + } + got := <-seen + if got.Get("X-Forwarded-For") != "198.51.100.7" { + t.Fatalf("X-Forwarded-For = %q", got.Get("X-Forwarded-For")) + } + if got.Get("Forwarded") != "" { + t.Fatalf("Forwarded was preserved: %q", got.Get("Forwarded")) + } + if got.Get("X-Forwarded-Proto") != "http" { + t.Fatalf("X-Forwarded-Proto = %q", got.Get("X-Forwarded-Proto")) + } + if got.Get("X-Forwarded-Host") != "" { + t.Fatalf("X-Forwarded-Host was forwarded: %q", got.Get("X-Forwarded-Host")) + } + for _, header := range []string{"X-Cxdb-Writer-Method", "X-Cxdb-Writer-Subject", "X-Cxdb-Writer-Issuer", "X-Cxdb-User-Email", "X-Cxdb-Unknown", "Authorization", "Cookie"} { + if got.Get(header) != "" { + t.Fatalf("%s was forwarded: %q", header, got.Get(header)) + } + } +} diff --git a/gateway/pkg/proxy/server.go b/gateway/pkg/proxy/server.go index 6b2cde1..c9b0236 100644 --- a/gateway/pkg/proxy/server.go +++ b/gateway/pkg/proxy/server.go @@ -5,17 +5,22 @@ package proxy import ( "context" + "encoding/json" "fmt" "io/fs" "log/slog" "net" "net/http" + "net/url" "strings" "sync" "time" + mcpauth "github.com/modelcontextprotocol/go-sdk/auth" + "github.com/modelcontextprotocol/go-sdk/oauthex" "github.com/strongdm/cxdb/gateway/internal/config" "github.com/strongdm/cxdb/gateway/pkg/auth" + "github.com/strongdm/cxdb/gateway/pkg/mcpserver" "golang.org/x/time/rate" ) @@ -25,14 +30,17 @@ type Server struct { mux *http.ServeMux sessions *auth.SessionStore google *auth.GoogleAuth + oidc *auth.BrowserOIDC + oauth *auth.OAuthServer proxy *ReverseProxy sse *SSEBroker logger *slog.Logger staticFS fs.FS - cspHeader string - hstsEnabled bool - limiters *ipRateLimiter + cspHeader string + hstsEnabled bool + limiters *ipRateLimiter + trustedProxyNetworks []*net.IPNet // Service-to-service auth verifiers (optional) tokenVerifiers []auth.BearerTokenVerifier @@ -74,6 +82,29 @@ func New(cfg config.Config, sessions *auth.SessionStore, google *auth.GoogleAuth hstsEnabled: strings.HasPrefix(strings.ToLower(cfg.PublicBaseURL), "https://"), limiters: newIPRateLimiter(rate.Limit(5), 10), } + for _, cidr := range cfg.TrustedProxyCIDRs { + _, network, err := net.ParseCIDR(cidr) + if err != nil { + return nil, fmt.Errorf("parse trusted proxy CIDR %q: %w", cidr, err) + } + s.trustedProxyNetworks = append(s.trustedProxyNetworks, network) + } + + if cfg.OIDCEnabled { + discoveryContext, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + browserOIDC, err := auth.NewBrowserOIDC(discoveryContext, cfg.OIDCIssuerURL, cfg.OIDCClientID, cfg.OIDCClientSecret, cfg.PublicBaseURL, cfg.OIDCAllowedDomains, sessions) + if err != nil { + return nil, fmt.Errorf("init browser OIDC: %w", err) + } + s.oidc = browserOIDC + } + oauthServer, err := auth.NewOAuthServer(sessions, cfg.PublicBaseURL, "/auth/login") + if err != nil { + return nil, fmt.Errorf("init OAuth server: %w", err) + } + s.oauth = oauthServer + s.tokenVerifiers = append(s.tokenVerifiers, auth.NewAPITokenVerifier(sessions), oauthServer) // Initialize K8s OIDC verifier if enabled if cfg.K8sOIDCEnabled { @@ -91,7 +122,7 @@ func New(cfg config.Config, sessions *auth.SessionStore, google *auth.GoogleAuth // Initialize AWS IAM token exchanger if enabled if cfg.AWSIAMEnabled { - // Extract issuer from PublicBaseURL (e.g., "https://your-domain.com" -> "your-domain.com") + // Extract issuer from PublicBaseURL (e.g., "https://cxdb.example.com" -> "cxdb.example.com") issuer := strings.TrimPrefix(cfg.PublicBaseURL, "https://") issuer = strings.TrimPrefix(issuer, "http://") issuer = strings.TrimSuffix(issuer, "/") @@ -115,9 +146,31 @@ func New(cfg config.Config, sessions *auth.SessionStore, google *auth.GoogleAuth mux.HandleFunc("/readyz", s.readyz) // OAuth endpoints (public) - mux.HandleFunc("/auth/google/login", google.LoginHandler) - mux.HandleFunc("/auth/google/callback", google.CallbackHandler) - mux.HandleFunc("/auth/google/logout", google.LogoutHandler) + mux.HandleFunc("/auth/login", s.login) + if google != nil { + mux.HandleFunc("/auth/google/login", google.LoginHandler) + mux.HandleFunc("/auth/google/callback", google.CallbackHandler) + mux.HandleFunc("/auth/google/logout", google.LogoutHandler) + } + if s.oidc != nil { + mux.HandleFunc("/auth/oidc/login", s.oidc.LoginHandler) + mux.HandleFunc("/auth/oidc/callback", s.oidc.CallbackHandler) + } + mux.HandleFunc("/.well-known/oauth-authorization-server", s.oauth.MetadataHandler) + mux.HandleFunc("/oauth/register", s.oauth.RegisterHandler) + mux.Handle("/oauth/authorize", http.NewCrossOriginProtection().Handler(http.HandlerFunc(s.oauth.AuthorizeHandler))) + mux.HandleFunc("/oauth/token", s.oauth.TokenHandler) + resourceMetadataURL := strings.TrimSuffix(cfg.PublicBaseURL, "/") + "/.well-known/oauth-protected-resource/mcp" + resourceMetadata := &oauthex.ProtectedResourceMetadata{ + Resource: strings.TrimSuffix(cfg.PublicBaseURL, "/") + "/mcp", AuthorizationServers: []string{strings.TrimSuffix(cfg.PublicBaseURL, "/")}, + ScopesSupported: []string{"cxdb:read", "cxdb:write"}, BearerMethodsSupported: []string{"header"}, ResourceName: "CXDB MCP", + } + mux.Handle("/.well-known/oauth-protected-resource/mcp", mcpauth.ProtectedResourceMetadataHandler(resourceMetadata)) + mcpHandler, err := mcpserver.New(cfg.CXDBBackendURL, resourceMetadataURL, s.tokenVerifiers, logger) + if err != nil { + return nil, fmt.Errorf("init MCP server: %w", err) + } + mux.Handle("/mcp", mcpHandler) // AWS IAM token exchange endpoint (public - uses AWS creds for auth) if s.awsExchanger != nil { @@ -126,6 +179,8 @@ func New(cfg config.Config, sessions *auth.SessionStore, google *auth.GoogleAuth // API info endpoint mux.HandleFunc("/api/v1/me", s.me) + mux.HandleFunc("/api/v1/tokens", s.tokens) + mux.HandleFunc("/api/v1/tokens/", s.tokenByID) // SSE endpoint for live events (must be before /v1/ catch-all) mux.Handle("/v1/events", sseBroker) @@ -145,14 +200,7 @@ func (s *Server) ListenAndServe(ctx context.Context) error { s.sse.Start(ctx) addr := fmt.Sprintf(":%s", s.cfg.Port) - handler := auth.RequireAuthForReadsWithOptions(auth.AuthMiddlewareOptions{ - Store: s.sessions, - DevBypass: s.cfg.DevMode, - TokenVerifiers: s.tokenVerifiers, - }, s.mux) - handler = s.rateLimitMiddleware(handler) - handler = s.securityHeaders(handler) - handler = s.loggingMiddleware(handler) + handler := s.Handler() srv := &http.Server{ Addr: addr, @@ -178,6 +226,25 @@ func (s *Server) ListenAndServe(ctx context.Context) error { return nil } +// Handler returns the complete production middleware stack. It is also used +// by integration tests so they exercise the same authorization boundary. +func (s *Server) Handler() http.Handler { + handler := auth.RequireAuthForReadsWithOptions(auth.AuthMiddlewareOptions{ + Store: s.sessions, + DevBypass: s.cfg.DevMode, + TokenVerifiers: s.tokenVerifiers, + }, s.mux) + handler = auth.RequireAuthForWrites(auth.AuthMiddlewareOptions{ + Store: s.sessions, + DevBypass: s.cfg.DevMode, + TokenVerifiers: s.tokenVerifiers, + }, handler) + handler = s.rateLimitMiddleware(handler) + handler = s.securityHeaders(handler) + handler = s.loggingMiddleware(handler) + return handler +} + func (s *Server) healthz(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) _, _ = w.Write([]byte("ok")) @@ -203,13 +270,133 @@ func (s *Server) me(w http.ResponseWriter, r *http.Request) { return } w.Header().Set("Content-Type", "application/json") - _, _ = fmt.Fprintf(w, `{"email":%q,"name":%q,"picture":%q}`, user.Email, user.Name, user.Picture) + writeJSONResponse(w, http.StatusOK, map[string]any{"email": user.Email, "name": user.Name, "picture": user.Picture, "issuer": user.Issuer, "subject": user.Subject, "scopes": user.Scopes, "auth_method": user.AuthMethod, "csrf_token": s.sessions.CSRFToken(user)}) +} + +func (s *Server) login(w http.ResponseWriter, r *http.Request) { + destination := "/auth/google/login" + if s.oidc != nil { + destination = "/auth/oidc/login" + } else if s.google == nil { + http.Error(w, "no browser login provider is configured", http.StatusServiceUnavailable) + return + } + if returnTo := r.URL.Query().Get("return_to"); returnTo != "" { + destination += "?return_to=" + url.QueryEscape(returnTo) + } + http.Redirect(w, r, destination, http.StatusFound) +} + +func (s *Server) tokens(w http.ResponseWriter, r *http.Request) { + user, ok := s.browserTokenManager(w, r) + if !ok { + return + } + switch r.Method { + case http.MethodGet: + tokens, err := s.sessions.ListPersonalAPITokens(r.Context(), user) + if err != nil { + http.Error(w, "unable to list tokens", http.StatusInternalServerError) + return + } + writeJSONResponse(w, http.StatusOK, map[string]any{"tokens": tokens}) + case http.MethodPost: + if !s.sessions.ValidCSRFToken(user, r.Header.Get("X-CSRF-Token")) { + http.Error(w, "invalid CSRF token", http.StatusForbidden) + return + } + r.Body = http.MaxBytesReader(w, r.Body, 16<<10) + var request struct { + Name string `json:"name"` + Scopes []string `json:"scopes"` + ExpiresAt string `json:"expires_at"` + } + if err := json.NewDecoder(r.Body).Decode(&request); err != nil { + http.Error(w, "invalid JSON", http.StatusBadRequest) + return + } + var expires time.Time + var err error + if request.ExpiresAt != "" { + expires, err = time.Parse(time.RFC3339, request.ExpiresAt) + if err != nil || expires.Before(time.Now()) { + http.Error(w, "expires_at must be a future RFC3339 time", http.StatusBadRequest) + return + } + } + metadata, plaintext, err := s.sessions.CreatePersonalAPIToken(r.Context(), user, request.Name, request.Scopes, expires) + if err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + writeJSONResponse(w, http.StatusCreated, map[string]any{"token": metadata, "plaintext": plaintext}) + default: + w.Header().Set("Allow", "GET, POST") + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + } +} + +func (s *Server) tokenByID(w http.ResponseWriter, r *http.Request) { + user, ok := s.browserTokenManager(w, r) + if !ok { + return + } + if r.Method != http.MethodDelete { + w.Header().Set("Allow", "DELETE") + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + return + } + if !s.sessions.ValidCSRFToken(user, r.Header.Get("X-CSRF-Token")) { + http.Error(w, "invalid CSRF token", http.StatusForbidden) + return + } + id := strings.TrimPrefix(r.URL.Path, "/api/v1/tokens/") + if id == "" || strings.Contains(id, "/") { + http.Error(w, "invalid token ID", http.StatusBadRequest) + return + } + if err := s.sessions.RevokePersonalAPIToken(r.Context(), user, id); err != nil { + http.Error(w, "token not found", http.StatusNotFound) + return + } + w.WriteHeader(http.StatusNoContent) +} + +func (s *Server) browserTokenManager(w http.ResponseWriter, r *http.Request) (*auth.Session, bool) { + if r.Header.Get("Authorization") != "" { + http.Error(w, "personal tokens cannot manage personal tokens", http.StatusForbidden) + return nil, false + } + user := auth.UserFromContext(r.Context()) + if user == nil { + user, _ = s.sessions.SessionFromRequest(r.Context(), r) + } + if user == nil || user.Subject == "" || user.Issuer == "" { + http.Error(w, "browser login required", http.StatusUnauthorized) + return nil, false + } + return user, true +} + +func writeJSONResponse(w http.ResponseWriter, status int, value any) { + w.Header().Set("Content-Type", "application/json") + w.Header().Set("Cache-Control", "no-store") + w.WriteHeader(status) + _ = json.NewEncoder(w).Encode(value) } // staticHandler serves the embedded React frontend with smart routing for Next.js static export. func (s *Server) staticHandler() http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { path := strings.TrimPrefix(r.URL.Path, "/") + switch path { + case "openapi.yaml": + w.Header().Set("Content-Type", "application/yaml") + case "openapi.json": + w.Header().Set("Content-Type", "application/json") + case "llms.txt": + w.Header().Set("Content-Type", "text/plain; charset=utf-8") + } // Handle root - serve index.html if path == "" { @@ -278,7 +465,7 @@ func (s *Server) rateLimitMiddleware(next http.Handler) http.Handler { next.ServeHTTP(w, r) return } - ip := clientIP(r) + ip := rateLimitClientIP(r, s.trustedProxyNetworks) limiter := s.limiters.get(ip) if !limiter.Allow() { s.logger.Warn("rate_limit_exceeded", "ip", ip, "path", r.URL.Path) @@ -294,7 +481,7 @@ func (s *Server) loggingMiddleware(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { // Skip wrapping for SSE endpoint - the wrapper can interfere with HTTP/2 streaming if r.URL.Path == "/v1/events" { - s.logger.Info("http_sse_start", "method", r.Method, "path", r.URL.Path, "ip", clientIP(r)) + s.logger.Info("http_sse_start", "method", r.Method, "path", r.URL.Path, "ip", observedPeerIP(r)) next.ServeHTTP(w, r) s.logger.Info("http_sse_end", "method", r.Method, "path", r.URL.Path) return @@ -315,7 +502,7 @@ func (s *Server) loggingMiddleware(next http.Handler) http.Handler { "status", sw.status, "duration_ms", time.Since(start).Milliseconds(), "size_bytes", sw.bytes, - "ip", clientIP(r), + "ip", observedPeerIP(r), "user", user, ) }) @@ -345,19 +532,6 @@ func (w *statusWriter) Flush() { } } -func clientIP(r *http.Request) string { - xff := r.Header.Get("X-Forwarded-For") - if xff != "" { - parts := strings.Split(xff, ",") - return strings.TrimSpace(parts[0]) - } - host, _, err := net.SplitHostPort(r.RemoteAddr) - if err != nil { - return r.RemoteAddr - } - return host -} - type ipRateLimiter struct { mu sync.Mutex visitors map[string]*rate.Limiter @@ -386,7 +560,7 @@ func (l *ipRateLimiter) get(ip string) *rate.Limiter { func shouldRateLimit(path string) bool { path = strings.ToLower(path) - if path == "/login" || strings.HasPrefix(path, "/auth/") { + if path == "/login" || strings.HasPrefix(path, "/auth/") || path == "/oauth/register" || path == "/oauth/token" || path == "/oauth/authorize" { return true } return false diff --git a/gateway/pkg/proxy/server_security_test.go b/gateway/pkg/proxy/server_security_test.go new file mode 100644 index 0000000..db8810a --- /dev/null +++ b/gateway/pkg/proxy/server_security_test.go @@ -0,0 +1,24 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + +package proxy + +import ( + "net/http" + "net/http/httptest" + "testing" +) + +func TestSecurityHeadersPreservesHandlerSpecificCSP(t *testing.T) { + server := &Server{cspHeader: "default-src 'self'; form-action 'self'"} + consentCSP := "default-src 'none'; form-action 'self' http://127.0.0.1:6276" + handler := server.securityHeaders(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Security-Policy", consentCSP) + w.WriteHeader(http.StatusOK) + })) + response := httptest.NewRecorder() + handler.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "/oauth/authorize", nil)) + if got := response.Header().Get("Content-Security-Policy"); got != consentCSP { + t.Fatalf("Content-Security-Policy = %q, want %q", got, consentCSP) + } +} diff --git a/gateway/pkg/proxy/server_tokens_test.go b/gateway/pkg/proxy/server_tokens_test.go new file mode 100644 index 0000000..7c9b7c9 --- /dev/null +++ b/gateway/pkg/proxy/server_tokens_test.go @@ -0,0 +1,193 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + +package proxy + +import ( + "bytes" + "encoding/json" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "path/filepath" + "strings" + "testing" + "testing/fstest" + "time" + + "github.com/strongdm/cxdb/gateway/internal/config" + "github.com/strongdm/cxdb/gateway/pkg/auth" +) + +func TestPersonalTokenHTTPCreateListRevokeAndCSRF(t *testing.T) { + store, err := auth.NewSessionStore(filepath.Join(t.TempDir(), "sessions.sqlite"), "session", time.Hour, "", false, "test-secret") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = store.Close() }) + sessionID, err := store.CreateForIdentity(t.Context(), "https://issuer.example", "alice", "alice@example.com", "Alice", "", "oidc", []string{"cxdb:read", "cxdb:write"}) + if err != nil { + t.Fatal(err) + } + session, err := store.Get(t.Context(), sessionID) + if err != nil { + t.Fatal(err) + } + server := &Server{sessions: store} + cookieRecorder := httptest.NewRecorder() + store.SetCookie(cookieRecorder, sessionID) + cookie := cookieRecorder.Result().Cookies()[0] + + badRequest := httptest.NewRequest(http.MethodPost, "/api/v1/tokens", strings.NewReader(`{"name":"laptop","scopes":["cxdb:read"]}`)) + badRequest.AddCookie(cookie) + badResponse := httptest.NewRecorder() + server.tokens(badResponse, badRequest) + if badResponse.Code != http.StatusForbidden { + t.Fatalf("missing CSRF status = %d", badResponse.Code) + } + + createRequest := httptest.NewRequest(http.MethodPost, "/api/v1/tokens", strings.NewReader(`{"name":"laptop","scopes":["cxdb:read","cxdb:write"]}`)) + createRequest.AddCookie(cookie) + createRequest.Header.Set("X-CSRF-Token", store.CSRFToken(session)) + createResponse := httptest.NewRecorder() + server.tokens(createResponse, createRequest) + if createResponse.Code != http.StatusCreated { + t.Fatalf("create status = %d, body=%s", createResponse.Code, createResponse.Body.String()) + } + var created struct { + Token auth.APIToken `json:"token"` + Plaintext string `json:"plaintext"` + } + if err := json.Unmarshal(createResponse.Body.Bytes(), &created); err != nil { + t.Fatal(err) + } + if created.Plaintext == "" { + t.Fatal("create did not return the one-time plaintext") + } + + listRequest := httptest.NewRequest(http.MethodGet, "/api/v1/tokens", nil) + listRequest.AddCookie(cookie) + listResponse := httptest.NewRecorder() + server.tokens(listResponse, listRequest) + if listResponse.Code != http.StatusOK || bytes.Contains(listResponse.Body.Bytes(), []byte(created.Plaintext)) || bytes.Contains(listResponse.Body.Bytes(), []byte("token_hash")) { + t.Fatalf("unsafe list response: status=%d body=%s", listResponse.Code, listResponse.Body.String()) + } + + bearerRequest := httptest.NewRequest(http.MethodGet, "/api/v1/tokens", nil) + bearerRequest.AddCookie(cookie) + bearerRequest.Header.Set("Authorization", "Bearer "+created.Plaintext) + bearerResponse := httptest.NewRecorder() + server.tokens(bearerResponse, bearerRequest) + if bearerResponse.Code != http.StatusForbidden { + t.Fatalf("bearer token management status = %d", bearerResponse.Code) + } + + revokeRequest := httptest.NewRequest(http.MethodDelete, "/api/v1/tokens/"+created.Token.ID, nil) + revokeRequest.AddCookie(cookie) + revokeRequest.Header.Set("X-CSRF-Token", store.CSRFToken(session)) + revokeResponse := httptest.NewRecorder() + server.tokenByID(revokeResponse, revokeRequest) + if revokeResponse.Code != http.StatusNoContent { + t.Fatalf("revoke status = %d", revokeResponse.Code) + } + if _, err := store.VerifyAPIToken(t.Context(), created.Plaintext); err == nil { + t.Fatal("revoked token still verifies") + } +} + +func TestProductionHandlerPersonalTokenLifecycleAndAPIUse(t *testing.T) { + store, err := auth.NewSessionStore(filepath.Join(t.TempDir(), "sessions.sqlite"), "session", time.Hour, "", false, "test-secret") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = store.Close() }) + backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"ok":true}`) + })) + t.Cleanup(backend.Close) + reverse, err := NewReverseProxy(backend.URL, slog.Default()) + if err != nil { + t.Fatal(err) + } + cfg := config.Config{ + PublicBaseURL: "http://localhost:8080", CXDBBackendURL: backend.URL, + Port: "0", DevMode: true, + } + server, err := New(cfg, store, nil, reverse, fstest.MapFS{ + "index.html": {Data: []byte("CXDB")}, + }, slog.Default()) + if err != nil { + t.Fatal(err) + } + remote := httptest.NewServer(server.Handler()) + t.Cleanup(remote.Close) + + meResponse, err := http.Get(remote.URL + "/api/v1/me") + if err != nil { + t.Fatal(err) + } + defer func() { _ = meResponse.Body.Close() }() + var me struct { + CSRFToken string `json:"csrf_token"` + } + if meResponse.StatusCode != http.StatusOK || json.NewDecoder(meResponse.Body).Decode(&me) != nil || me.CSRFToken == "" { + t.Fatalf("me response status = %d", meResponse.StatusCode) + } + + createRequest, _ := http.NewRequest(http.MethodPost, remote.URL+"/api/v1/tokens", strings.NewReader(`{"name":"integration","scopes":["cxdb:read","cxdb:write"]}`)) + createRequest.Header.Set("Content-Type", "application/json") + createRequest.Header.Set("X-CSRF-Token", me.CSRFToken) + createResponse, err := http.DefaultClient.Do(createRequest) + if err != nil { + t.Fatal(err) + } + defer func() { _ = createResponse.Body.Close() }() + var created struct { + Token auth.APIToken `json:"token"` + Plaintext string `json:"plaintext"` + } + if createResponse.StatusCode != http.StatusCreated || json.NewDecoder(createResponse.Body).Decode(&created) != nil || created.Plaintext == "" { + t.Fatalf("create response status = %d", createResponse.StatusCode) + } + + apiRequest, _ := http.NewRequest(http.MethodGet, remote.URL+"/v1/private", nil) + apiRequest.Header.Set("Authorization", "Bearer "+created.Plaintext) + apiResponse, err := http.DefaultClient.Do(apiRequest) + if err != nil { + t.Fatal(err) + } + if err := apiResponse.Body.Close(); err != nil { + t.Fatal(err) + } + if apiResponse.StatusCode != http.StatusOK { + t.Fatalf("API token request status = %d", apiResponse.StatusCode) + } + + revokeRequest, _ := http.NewRequest(http.MethodDelete, remote.URL+"/api/v1/tokens/"+created.Token.ID, nil) + revokeRequest.Header.Set("X-CSRF-Token", me.CSRFToken) + revokeResponse, err := http.DefaultClient.Do(revokeRequest) + if err != nil { + t.Fatal(err) + } + if err := revokeResponse.Body.Close(); err != nil { + t.Fatal(err) + } + if revokeResponse.StatusCode != http.StatusNoContent { + t.Fatalf("revoke response status = %d", revokeResponse.StatusCode) + } + + revokedRequest, _ := http.NewRequest(http.MethodGet, remote.URL+"/v1/private", nil) + revokedRequest.Header.Set("Authorization", "Bearer "+created.Plaintext) + revokedResponse, err := http.DefaultClient.Do(revokedRequest) + if err != nil { + t.Fatal(err) + } + if err := revokedResponse.Body.Close(); err != nil { + t.Fatal(err) + } + if revokedResponse.StatusCode != http.StatusUnauthorized { + t.Fatalf("revoked token status = %d", revokedResponse.StatusCode) + } +} diff --git a/server/Cargo.toml b/server/Cargo.toml index 1e7b3d4..aec9eb4 100644 --- a/server/Cargo.toml +++ b/server/Cargo.toml @@ -31,6 +31,7 @@ sysinfo = "0.30" libc = "0.2" regex = "1.10" tracing = "0.1" +rayon = "1.10" # AWS SDK for S3 sync (optional feature for production deployments) aws-config = { version = "1.5", features = ["behavior-version-latest"] } diff --git a/server/src/blob_store/mod.rs b/server/src/blob_store/mod.rs index 1ab1dd3..819e24d 100644 --- a/server/src/blob_store/mod.rs +++ b/server/src/blob_store/mod.rs @@ -1,15 +1,16 @@ // Copyright 2025 StrongDM Inc // SPDX-License-Identifier: Apache-2.0 -use std::collections::HashMap; +use std::collections::{HashMap, HashSet}; use std::fs::{File, OpenOptions}; -use std::io::{Read, Seek, SeekFrom, Write}; +use std::io::{Cursor, Read, Seek, SeekFrom, Write}; #[cfg(unix)] use std::os::unix::fs::FileExt; use std::path::{Path, PathBuf}; use byteorder::{LittleEndian, ReadBytesExt, WriteBytesExt}; use crc32fast::Hasher; +use rayon::prelude::*; use crate::error::{Result, StoreError}; @@ -296,10 +297,110 @@ impl BlobStore { if raw_bytes.len() as u32 != raw_len { return Err(StoreError::Corrupt("blob length mismatch".into())); } + if blake3::hash(&raw_bytes).as_bytes() != hash { + return Err(StoreError::Corrupt("blob content hash mismatch".into())); + } Ok(raw_bytes) } + /// Read several blobs with bounded, coalesced pread operations. + /// + /// Records are decoded independently after the range read. The returned + /// vector has exactly the same order and duplicates as `hashes`. + pub fn get_many(&self, hashes: &[[u8; 32]]) -> Result>> { + const HEADER_SIZE: u64 = 48; + const CRC_SIZE: u64 = 4; + const MAX_GAP: u64 = 64 * 1024; + const MAX_RANGE: u64 = 16 * 1024 * 1024; + + if hashes.is_empty() { + return Ok(Vec::new()); + } + let mut unique = Vec::with_capacity(hashes.len()); + let mut seen = HashSet::with_capacity(hashes.len()); + for hash in hashes { + if seen.insert(*hash) { + let entry = self + .index + .get(hash) + .ok_or_else(|| StoreError::NotFound("blob".into()))? + .clone(); + let end = entry + .offset + .checked_add(HEADER_SIZE) + .and_then(|v| v.checked_add(u64::from(entry.stored_len))) + .and_then(|v| v.checked_add(CRC_SIZE)) + .ok_or_else(|| StoreError::Corrupt("blob record offset overflow".into()))?; + unique.push((*hash, entry, end)); + } + } + unique.sort_unstable_by_key(|(_, entry, _)| entry.offset); + + let pack_len = self.pack_read.metadata()?.len(); + let mut decoded = HashMap::with_capacity(unique.len()); + let mut first = 0; + while first < unique.len() { + let range_start = unique[first].1.offset; + let mut range_end = unique[first].2; + let mut last = first + 1; + while last < unique.len() { + let next = &unique[last]; + if next.1.offset < range_end { + return Err(StoreError::Corrupt("overlapping blob index entries".into())); + } + let gap = next.1.offset - range_end; + let span = next + .2 + .checked_sub(range_start) + .ok_or_else(|| StoreError::Corrupt("invalid blob index range".into()))?; + if gap > MAX_GAP || span > MAX_RANGE { + break; + } + range_end = next.2; + last += 1; + } + if range_end > pack_len { + return Err(StoreError::Corrupt( + "blob index points past pack end".into(), + )); + } + let range_len = usize::try_from(range_end - range_start) + .map_err(|_| StoreError::Corrupt("blob read range exceeds address space".into()))?; + let mut range = vec![0u8; range_len]; + self.read_at_exact(range_start, &mut range)?; + let group: Result)>> = unique[first..last] + .par_iter() + .map(|(hash, entry, end)| { + let start = usize::try_from(entry.offset - range_start).map_err(|_| { + StoreError::Corrupt("blob offset exceeds address space".into()) + })?; + let end = usize::try_from(*end - range_start).map_err(|_| { + StoreError::Corrupt("blob record exceeds address space".into()) + })?; + let record = range.get(start..end).ok_or_else(|| { + StoreError::Corrupt("blob record outside read range".into()) + })?; + Ok((*hash, decode_blob_record_slice(record, hash, entry)?)) + }) + .collect(); + for (hash, payload) in group? { + decoded.insert(hash, payload); + } + first = last; + } + + hashes + .iter() + .map(|hash| { + decoded + .get(hash) + .cloned() + .ok_or_else(|| StoreError::NotFound("blob".into())) + }) + .collect() + } + /// Read exactly buf.len() bytes from the read handle at the given offset using pread. fn read_at_exact(&self, offset: u64, buf: &mut [u8]) -> Result<()> { let mut total_read = 0usize; @@ -335,6 +436,59 @@ impl BlobStore { } } +fn decode_blob_record_slice( + record: &[u8], + expected_hash: &[u8; 32], + expected_entry: &BlobIndexEntry, +) -> Result> { + const HEADER_SIZE: usize = 48; + let expected_len = HEADER_SIZE + .checked_add(expected_entry.stored_len as usize) + .and_then(|v| v.checked_add(4)) + .ok_or_else(|| StoreError::Corrupt("blob record length overflow".into()))?; + if record.len() != expected_len { + return Err(StoreError::Corrupt("blob record length mismatch".into())); + } + let mut header = Cursor::new(&record[..HEADER_SIZE]); + let magic = header.read_u32::()?; + let version = header.read_u16::()?; + let codec_raw = header.read_u16::()?; + let raw_len = header.read_u32::()?; + let stored_len = header.read_u32::()?; + let mut stored_hash = [0u8; 32]; + header.read_exact(&mut stored_hash)?; + if magic != BLOB_MAGIC || version != BLOB_VERSION { + return Err(StoreError::Corrupt("invalid blob header".into())); + } + if &stored_hash != expected_hash + || raw_len != expected_entry.raw_len + || stored_len != expected_entry.stored_len + || codec_raw != expected_entry.codec as u16 + { + return Err(StoreError::Corrupt("blob index/header mismatch".into())); + } + let stored_end = HEADER_SIZE + stored_len as usize; + let stored = &record[HEADER_SIZE..stored_end]; + let crc = Cursor::new(&record[stored_end..]).read_u32::()?; + let mut hasher = Hasher::new(); + hasher.update(&record[..stored_end]); + if crc != hasher.finalize() { + return Err(StoreError::Corrupt("blob crc mismatch".into())); + } + let raw = match expected_entry.codec { + BlobCodec::None => stored.to_vec(), + BlobCodec::Zstd => zstd::decode_all(stored) + .map_err(|e| StoreError::Corrupt(format!("zstd decode failed: {e}")))?, + }; + if raw.len() != raw_len as usize { + return Err(StoreError::Corrupt("blob length mismatch".into())); + } + if blake3::hash(&raw).as_bytes() != expected_hash { + return Err(StoreError::Corrupt("blob content hash mismatch".into())); + } + Ok(raw) +} + #[derive(Debug, Clone)] pub struct BlobStoreStats { pub blobs_total: usize, @@ -345,3 +499,28 @@ pub struct BlobStoreStats { fn file_len(path: &PathBuf) -> u64 { std::fs::metadata(path).map(|m| m.len()).unwrap_or(0) } + +#[cfg(test)] +mod tests { + use super::BlobStore; + use tempfile::tempdir; + + #[test] + fn get_many_preserves_order_and_duplicates() { + let dir = tempdir().expect("tempdir"); + let mut store = BlobStore::open(dir.path()).expect("open"); + let first = b"first payload"; + let second = b"second payload"; + let first_hash = *blake3::hash(first).as_bytes(); + let second_hash = *blake3::hash(second).as_bytes(); + store.put_if_absent(first_hash, first).expect("first"); + store.put_if_absent(second_hash, second).expect("second"); + let values = store + .get_many(&[second_hash, first_hash, second_hash]) + .expect("batch read"); + assert_eq!( + values, + vec![second.to_vec(), first.to_vec(), second.to_vec()] + ); + } +} diff --git a/server/src/http/mod.rs b/server/src/http/mod.rs index 91a304b..b1f4bde 100644 --- a/server/src/http/mod.rs +++ b/server/src/http/mod.rs @@ -2,7 +2,7 @@ // SPDX-License-Identifier: Apache-2.0 use std::collections::HashMap; -use std::io::Write; +use std::io::{Read, Write}; use std::sync::{Arc, Mutex, RwLock}; use std::thread; use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; @@ -17,13 +17,20 @@ use crate::error::{Result, StoreError}; use crate::events::{EventBus, StoreEvent}; use crate::fs_store::EntryKind; use crate::metrics::{Metrics, SessionTracker}; -use crate::projection::{BytesRender, EnumRender, RenderOptions, TimeRender, U64Format}; +use crate::projection::{ + assemble_turn_page_json, serialize_turn_page, BytesRender, EnumRender, RenderOptions, + TimeRender, TurnProjectionOptions, U64Format, +}; use crate::registry::{ FieldSpec, ItemsSpec, PutOutcome, Registry, RegistryBundle, RendererSpec, TypeVersionSpec, }; use crate::store::Store; type HttpResponse = (u16, Response>>); +const MAX_HTTP_PAYLOAD_BYTES: usize = 4 * 1024 * 1024; +const MAX_HTTP_JSON_BODY_BYTES: usize = 24 * 1024 * 1024; +const MAX_APPEND_BATCH_ITEMS: usize = 256; +const MAX_APPEND_BATCH_PAYLOAD_BYTES: usize = 16 * 1024 * 1024; pub fn start_http( bind_addr: String, @@ -660,14 +667,21 @@ fn handle_request( let type_id = get_required_string(&body, "type_id")?; let type_version = get_required_u32(&body, "type_version")?; let parent_turn_id = get_optional_u64(&body, "parent_turn_id")?.unwrap_or(0); - let payload_json = body - .get("data") - .or_else(|| body.get("payload")) - .ok_or_else(|| { - StoreError::InvalidInput("missing required field: data or payload".into()) - })?; - - let payload_bytes = { + let payload_bytes = if let Some(encoded) = + body.get("payload_base64").and_then(JsonValue::as_str) + { + let payload = decode_http_payload_base64(encoded)?; + validate_http_payload(&payload)?; + payload + } else { + let payload_json = body + .get("data") + .or_else(|| body.get("payload")) + .ok_or_else(|| { + StoreError::InvalidInput( + "missing required field: payload_base64, data, or payload".into(), + ) + })?; let registry = registry.lock().unwrap(); encode_http_payload(payload_json, &type_id, type_version, ®istry)? }; @@ -736,6 +750,137 @@ fn handle_request( ), )) } + (Method::Post, ["v1", "contexts", context_id, "append-batch"]) => { + let context_id: u64 = context_id + .parse() + .map_err(|_| StoreError::InvalidInput("invalid context_id".into()))?; + let body = parse_json_body(&mut request)?; + let items = body + .get("turns") + .or_else(|| body.get("items")) + .and_then(JsonValue::as_array) + .ok_or_else(|| { + StoreError::InvalidInput("missing required field: turns".into()) + })?; + if items.is_empty() { + return Err(StoreError::InvalidInput( + "append batch must contain at least one turn".into(), + )); + } + if items.len() > MAX_APPEND_BATCH_ITEMS { + return Err(StoreError::InvalidInput(format!( + "append batch must contain at most {MAX_APPEND_BATCH_ITEMS} turns" + ))); + } + let registry_guard = registry.lock().unwrap(); + let mut prepared = Vec::with_capacity(items.len()); + let mut total_payload_bytes = 0usize; + for (index, item) in items.iter().enumerate() { + match prepare_batch_item(item, ®istry_guard) { + Ok(prepared_item) => { + total_payload_bytes = total_payload_bytes + .checked_add(prepared_item.2.len()) + .ok_or_else(|| { + StoreError::InvalidInput( + "append batch payload size overflow".into(), + ) + })?; + if total_payload_bytes > MAX_APPEND_BATCH_PAYLOAD_BYTES { + return Err(StoreError::InvalidInput( + "append batch payloads must total at most 16 MiB".into(), + )); + } + prepared.push(prepared_item) + } + Err(error) => { + return batch_failure_response(context_id, &[], index, &error) + } + } + } + drop(registry_guard); + + // Hold the writer lock for the complete batch. This prevents an + // unrelated append from being inserted between two batch items. + // Storage is append-only, so a later I/O or parent error cannot + // roll back earlier records. Such a failure returns the exact + // successful prefix and failed index. + let mut store_guard = store.write().unwrap(); + let mut appended = Vec::with_capacity(prepared.len()); + let mut current_head = store_guard.get_head(context_id)?.head_turn_id; + for (index, (type_id, type_version, payload, requested_parent)) in + prepared.into_iter().enumerate() + { + let parent = requested_parent.unwrap_or(current_head); + let hash = blake3::hash(&payload); + let (record, metadata) = match store_guard.append_turn( + context_id, + parent, + type_id.clone(), + type_version, + 1, + 0, + payload.len() as u32, + *hash.as_bytes(), + &payload, + ) { + Ok(result) => result, + Err(error) => { + drop(store_guard); + return batch_failure_response(context_id, &appended, index, &error); + } + }; + current_head = record.turn_id; + event_bus.publish(StoreEvent::TurnAppended { + context_id: context_id.to_string(), + turn_id: record.turn_id.to_string(), + parent_turn_id: record.parent_turn_id.to_string(), + depth: record.depth, + declared_type_id: Some(type_id), + declared_type_version: Some(type_version), + }); + if let Some(meta) = metadata { + event_bus.publish(StoreEvent::ContextMetadataUpdated { + context_id: context_id.to_string(), + client_tag: meta.client_tag, + title: meta.title, + labels: meta.labels, + has_provenance: meta.provenance.is_some(), + }); + if let Some(prov) = meta.provenance { + if let Some(parent_context_id) = prov.parent_context_id { + event_bus.publish(StoreEvent::ContextLinked { + child_context_id: context_id.to_string(), + parent_context_id: parent_context_id.to_string(), + root_context_id: prov.root_context_id.map(|v| v.to_string()), + spawn_reason: prov.spawn_reason, + }); + } + } + } + appended.push(json!({ + "turn_id": format_id(record.turn_id, &U64Format::Number), + "parent_turn_id": format_id(record.parent_turn_id, &U64Format::Number), + "depth": record.depth, + "content_hash": hex::encode(hash.as_bytes()), + })); + } + drop(store_guard); + let bytes = serde_json::to_vec(&json!({ + "context_id": format_id(context_id, &U64Format::Number), + "turns": appended, + "partial": false, + })) + .map_err(|e| StoreError::InvalidInput(format!("json encode error: {e}")))?; + Ok(( + 201, + Response::from_data(bytes) + .with_status_code(StatusCode(201)) + .with_header( + Header::from_bytes(&b"Content-Type"[..], &b"application/json"[..]) + .unwrap(), + ), + )) + } (Method::Get, ["v1", "contexts", context_id, "turns"]) => { let context_id: u64 = context_id .parse() @@ -747,8 +892,21 @@ fn handle_request( .unwrap_or(64); let before_turn_id = params .get("before_turn_id") - .and_then(|v| v.parse::().ok()) + .map(|value| { + value + .parse::() + .map_err(|_| StoreError::InvalidInput("invalid before_turn_id".into())) + }) + .transpose()? .unwrap_or(0); + let exact_turn_id = params + .get("turn_id") + .map(|value| { + value + .parse::() + .map_err(|_| StoreError::InvalidInput("invalid turn_id".into())) + }) + .transpose()?; let view = params.get("view").map(|v| v.as_str()).unwrap_or("typed"); let type_hint_mode = params .get("type_hint_mode") @@ -777,6 +935,20 @@ fn handle_request( .get("include_unknown") .map(|v| v == "1") .unwrap_or(false); + let string_limit = match params.get("string_limit") { + Some(value) => { + let value = value + .parse::() + .map_err(|_| StoreError::InvalidInput("invalid string_limit".into()))?; + if value > 64 * 1024 { + return Err(StoreError::InvalidInput( + "string_limit exceeds 65536".into(), + )); + } + Some(value) + } + None => None, + }; let as_type_id = params.get("as_type_id").cloned(); let as_type_version = params @@ -789,12 +961,19 @@ fn handle_request( enum_render, time_render, include_unknown, + string_limit: if exact_turn_id.is_some() { + None + } else { + string_limit + }, }; let store = store.read().unwrap(); let head = store.get_head(context_id)?; let t0 = Instant::now(); - let turns = if before_turn_id == 0 { + let turns = if let Some(turn_id) = exact_turn_id { + vec![store.get_context_turn(context_id, turn_id, true)?] + } else if before_turn_id == 0 { store.get_last(context_id, limit, true)? } else { store.get_before(context_id, before_turn_id, limit, true)? @@ -802,135 +981,33 @@ fn handle_request( metrics.record_get_last(t0.elapsed()); let registry = registry.lock().unwrap(); - let mut out_turns = Vec::new(); - for item in turns.iter() { - let declared_type_id = item.meta.declared_type_id.clone(); - let declared_type_version = item.meta.declared_type_version; - - let (decoded_type_id, decoded_type_version) = match type_hint_mode { - "explicit" => { - let id = as_type_id.clone().ok_or_else(|| { - StoreError::InvalidInput("as_type_id required".into()) - })?; - let ver = as_type_version.ok_or_else(|| { - StoreError::InvalidInput("as_type_version required".into()) - })?; - (id, ver) - } - "latest" => { - let latest = registry - .get_latest_type_version(&declared_type_id) - .ok_or_else(|| StoreError::NotFound("type descriptor".into()))?; - (declared_type_id.clone(), latest.version) - } - _ => (declared_type_id.clone(), declared_type_version), - }; - - let mut turn_obj = Map::new(); - turn_obj.insert( - "turn_id".into(), - format_id(item.record.turn_id, &u64_format), - ); - turn_obj.insert( - "parent_turn_id".into(), - format_id(item.record.parent_turn_id, &u64_format), - ); - turn_obj.insert("depth".into(), JsonValue::Number(item.record.depth.into())); - turn_obj.insert( - "declared_type".into(), - json!({ - "type_id": declared_type_id, - "type_version": declared_type_version, - }), - ); - - if view == "typed" || view == "both" { - let desc = registry - .get_type_version(&decoded_type_id, decoded_type_version) - .ok_or_else(|| StoreError::NotFound("type descriptor".into()))?; - let payload = item - .payload - .as_ref() - .ok_or_else(|| StoreError::InvalidInput("payload not loaded".into()))?; - let projected = - crate::projection::project_msgpack(payload, desc, ®istry, &options)?; - turn_obj.insert( - "decoded_as".into(), - json!({ - "type_id": decoded_type_id, - "type_version": decoded_type_version, - }), - ); - turn_obj.insert("data".into(), projected.data); - if let Some(unknown) = projected.unknown { - turn_obj.insert("unknown".into(), unknown); - } - } - - if view == "raw" || view == "both" { - let raw_payload = item - .payload - .as_ref() - .ok_or_else(|| StoreError::InvalidInput("payload not loaded".into()))?; - turn_obj.insert( - "content_hash_b3".into(), - JsonValue::String(hex::encode(item.record.payload_hash)), - ); - turn_obj.insert( - "encoding".into(), - JsonValue::Number(item.meta.encoding.into()), - ); - turn_obj.insert("compression".into(), JsonValue::Number(0u32.into())); - turn_obj.insert( - "uncompressed_len".into(), - JsonValue::Number((raw_payload.len() as u32).into()), - ); - match bytes_render { - BytesRender::Base64 => { - turn_obj.insert( - "bytes_b64".into(), - JsonValue::String( - base64::engine::general_purpose::STANDARD - .encode(raw_payload), - ), - ); - } - BytesRender::Hex => { - turn_obj.insert( - "bytes_hex".into(), - JsonValue::String(hex::encode(raw_payload)), - ); - } - BytesRender::LenOnly => { - turn_obj.insert( - "bytes_len".into(), - JsonValue::Number((raw_payload.len() as u64).into()), - ); - } - } - } - - out_turns.push(JsonValue::Object(turn_obj)); - } - - let next_before = turns - .first() - .map(|t| format_id(t.record.turn_id, &u64_format)); - let meta = json!({ + let projection_options = TurnProjectionOptions { + view, + type_hint_mode, + as_type_id: as_type_id.as_deref(), + as_type_version, + render: &options, + }; + let serialized_turns = serialize_turn_page(&turns, ®istry, &projection_options)?; + let next_before = if exact_turn_id.is_some() { + None + } else { + turns + .first() + .map(|t| format_id(t.record.turn_id, &u64_format)) + }; + let mut meta = json!({ "context_id": format_id(context_id, &u64_format), "head_turn_id": format_id(head.head_turn_id, &u64_format), "head_depth": head.head_depth, "registry_bundle_id": registry.last_bundle_id(), }); - - let resp = json!({ - "meta": meta, - "turns": out_turns, - "next_before_turn_id": next_before, - }); - - let bytes = serde_json::to_vec(&resp) - .map_err(|e| StoreError::InvalidInput(format!("json encode error: {e}")))?; + if exact_turn_id.is_none() { + if let Some(string_limit) = string_limit { + meta["string_limit"] = json!(string_limit); + } + } + let bytes = assemble_turn_page_json(&meta, next_before, &serialized_turns)?; Ok(( 200, Response::from_data(bytes) @@ -1415,7 +1492,15 @@ fn parse_base_turn_id( fn parse_json_body(request: &mut tiny_http::Request) -> Result { let mut body = Vec::new(); - request.as_reader().read_to_end(&mut body)?; + request + .as_reader() + .take((MAX_HTTP_JSON_BODY_BYTES + 1) as u64) + .read_to_end(&mut body)?; + if body.len() > MAX_HTTP_JSON_BODY_BYTES { + return Err(StoreError::InvalidInput( + "json request body must be at most 24 MiB".into(), + )); + } if body.is_empty() { return Ok(JsonValue::Object(Map::new())); } @@ -1442,6 +1527,86 @@ fn parse_json_u64(value: &JsonValue, field_name: &str) -> Result { } } +fn decode_http_payload_base64(encoded: &str) -> Result> { + let payload = base64::engine::general_purpose::STANDARD + .decode(encoded) + .map_err(|e| StoreError::InvalidInput(format!("invalid payload_base64: {e}")))?; + if payload.len() > MAX_HTTP_PAYLOAD_BYTES { + return Err(StoreError::InvalidInput( + "payload_base64 must contain at most 4 MiB".into(), + )); + } + Ok(payload) +} + +fn validate_http_payload(payload: &[u8]) -> Result<()> { + let mut cursor = std::io::Cursor::new(payload); + let value = rmpv::decode::read_value(&mut cursor) + .map_err(|e| StoreError::InvalidInput(format!("msgpack decode error: {e}")))?; + if cursor.position() != payload.len() as u64 { + return Err(StoreError::InvalidInput( + "payload contains trailing msgpack data".into(), + )); + } + if !matches!(value, MsgpackValue::Map(_)) { + return Err(StoreError::InvalidInput("payload is not a map".into())); + } + Ok(()) +} + +fn prepare_batch_item( + item: &JsonValue, + registry: &Registry, +) -> Result<(String, u32, Vec, Option)> { + let type_id = get_required_string(item, "type_id")?; + let type_version = get_required_u32(item, "type_version")?; + let payload = if let Some(encoded) = item.get("payload_base64").and_then(JsonValue::as_str) { + let payload = decode_http_payload_base64(encoded)?; + validate_http_payload(&payload)?; + payload + } else { + let value = item + .get("data") + .or_else(|| item.get("payload")) + .ok_or_else(|| { + StoreError::InvalidInput( + "missing required field: payload_base64, data, or payload".into(), + ) + })?; + encode_http_payload(value, &type_id, type_version, registry)? + }; + let parent = get_optional_u64(item, "parent_turn_id")?; + Ok((type_id, type_version, payload, parent)) +} + +fn batch_failure_response( + context_id: u64, + appended: &[JsonValue], + failed_index: usize, + error: &StoreError, +) -> Result { + let (status, message) = map_error(error); + let bytes = serde_json::to_vec(&json!({ + "context_id": format_id(context_id, &U64Format::Number), + "turns": appended, + "partial": !appended.is_empty(), + "failed_index": failed_index, + "error": { + "code": status, + "message": message, + }, + })) + .map_err(|e| StoreError::InvalidInput(format!("json encode error: {e}")))?; + Ok(( + status, + Response::from_data(bytes) + .with_status_code(StatusCode(status)) + .with_header( + Header::from_bytes(&b"Content-Type"[..], &b"application/json"[..]).unwrap(), + ), + )) +} + fn get_required_string(body: &JsonValue, key: &str) -> Result { body.get(key) .and_then(|v| v.as_str()) @@ -1968,4 +2133,68 @@ mod tests { *k == MsgpackValue::from("text") && *v == MsgpackValue::from("hello") })); } + + #[test] + fn decode_http_payload_base64_accepts_mcp_payloads() { + let encoded = base64::engine::general_purpose::STANDARD.encode(b"mcp payload"); + assert_eq!( + decode_http_payload_base64(&encoded).expect("decode payload"), + b"mcp payload" + ); + } + + #[test] + fn decode_http_payload_base64_rejects_oversized_payloads() { + let encoded = + base64::engine::general_purpose::STANDARD + .encode(vec![0_u8; MAX_HTTP_PAYLOAD_BYTES + 1]); + let error = decode_http_payload_base64(&encoded).expect_err("oversized payload"); + assert!(error.to_string().contains("at most 4 MiB")); + } + + #[test] + fn validate_http_payload_requires_complete_msgpack_map() { + let mut valid = Vec::new(); + rmpv::encode::write_value( + &mut valid, + &MsgpackValue::Map(vec![( + MsgpackValue::Integer(1.into()), + MsgpackValue::String("hello".into()), + )]), + ) + .expect("encode map"); + validate_http_payload(&valid).expect("valid payload"); + + let mut trailing = valid.clone(); + trailing.push(0xc0); + assert!(validate_http_payload(&trailing) + .expect_err("trailing data") + .to_string() + .contains("trailing")); + assert!(validate_http_payload(&[0xc0]) + .expect_err("non-map") + .to_string() + .contains("not a map")); + } + + #[test] + fn batch_failures_preserve_route_status_semantics() { + let (conflict_status, _) = batch_failure_response( + 1, + &[json!({"turn_id": 1})], + 1, + &StoreError::NotFound("parent turn".into()), + ) + .expect("conflict response"); + assert_eq!(conflict_status, 409); + + let (storage_status, _) = batch_failure_response( + 1, + &[], + 0, + &StoreError::Corrupt("blob content hash mismatch".into()), + ) + .expect("storage response"); + assert_eq!(storage_status, 500); + } } diff --git a/server/src/projection/fast.rs b/server/src/projection/fast.rs new file mode 100644 index 0000000..7e11563 --- /dev/null +++ b/server/src/projection/fast.rs @@ -0,0 +1,31 @@ +// Copyright 2025 StrongDM Inc +// SPDX-License-Identifier: Apache-2.0 + +//! Streaming page serializer entry point. +//! +//! The page scheduler owns parallelism. This serializer keeps the wire shape +//! in one place and delegates field semantics to the compatibility projector, +//! so named and numeric key handling cannot drift between paths. + +use serde::Serialize; + +use super::{project_turn, TurnProjectionOptions}; +use crate::error::{Result, StoreError}; +use crate::registry::Registry; +use crate::store::TurnWithMeta; + +pub(super) fn serialize_turn( + item: &TurnWithMeta, + registry: &Registry, + options: &TurnProjectionOptions<'_>, +) -> Result> { + let projected = project_turn(item, registry, options)?; + serde_json::to_vec(&projected) + .map_err(|error| StoreError::InvalidInput(format!("json encode error: {error}"))) +} + +#[allow(dead_code)] +fn _serialize_json(value: &T) -> Result> { + serde_json::to_vec(value) + .map_err(|error| StoreError::InvalidInput(format!("json encode error: {error}"))) +} diff --git a/server/src/projection/mod.rs b/server/src/projection/mod.rs index 36997ad..0a1965f 100644 --- a/server/src/projection/mod.rs +++ b/server/src/projection/mod.rs @@ -1,6 +1,8 @@ // Copyright 2025 StrongDM Inc // SPDX-License-Identifier: Apache-2.0 +use rayon::prelude::*; +use std::borrow::Cow; use std::collections::HashMap; use base64::Engine; @@ -44,13 +46,223 @@ pub struct RenderOptions { pub enum_render: EnumRender, pub time_render: TimeRender, pub include_unknown: bool, + pub string_limit: Option, } +pub struct TurnProjectionOptions<'a> { + pub view: &'a str, + pub type_hint_mode: &'a str, + pub as_type_id: Option<&'a str>, + pub as_type_version: Option, + pub render: &'a RenderOptions, +} + +pub fn project_turn_page( + turns: &[crate::store::TurnWithMeta], + registry: &Registry, + options: &TurnProjectionOptions<'_>, +) -> Result> { + if turns.len() < 8 { + turns + .iter() + .map(|turn| project_turn(turn, registry, options)) + .collect() + } else { + turns + .par_iter() + .map(|turn| project_turn(turn, registry, options)) + .collect() + } +} + +pub fn serialize_turn_page( + turns: &[crate::store::TurnWithMeta], + registry: &Registry, + options: &TurnProjectionOptions<'_>, +) -> Result>> { + if turns.len() < 8 { + turns + .iter() + .map(|turn| fast::serialize_turn(turn, registry, options)) + .collect() + } else { + turns + .par_iter() + .map(|turn| fast::serialize_turn(turn, registry, options)) + .collect() + } +} + +pub fn assemble_turn_page_json( + meta: &JsonValue, + next_before_turn_id: Option, + turns: &[Vec], +) -> Result> { + let meta = serde_json::to_vec(meta) + .map_err(|e| StoreError::InvalidInput(format!("json encode error: {e}")))?; + let next = serde_json::to_vec(&next_before_turn_id) + .map_err(|e| StoreError::InvalidInput(format!("json encode error: {e}")))?; + let mut out = Vec::with_capacity(meta.len() + turns.iter().map(Vec::len).sum::() + 64); + out.extend_from_slice(br#"{"meta":"#); + out.extend_from_slice(&meta); + out.extend_from_slice(br#","turns":["#); + for (index, turn) in turns.iter().enumerate() { + if index != 0 { + out.push(b','); + } + out.extend_from_slice(turn); + } + out.extend_from_slice(br#"],"next_before_turn_id":"#); + out.extend_from_slice(&next); + out.push(b'}'); + Ok(out) +} + +#[cfg(test)] +mod page_assembly_tests { + use super::assemble_turn_page_json; + use serde_json::json; + + #[test] + fn assembles_empty_and_populated_pages_as_json() { + let empty = assemble_turn_page_json(&json!({"context_id": 1}), None, &[]) + .expect("assemble empty page"); + assert_eq!( + serde_json::from_slice::(&empty).expect("parse empty page"), + json!({"meta": {"context_id": 1}, "turns": [], "next_before_turn_id": null}) + ); + + let turns = vec![br#"{"turn_id":1}"#.to_vec(), br#"{"turn_id":2}"#.to_vec()]; + let populated = assemble_turn_page_json(&json!({"context_id": 1}), Some(json!(2)), &turns) + .expect("assemble populated page"); + assert_eq!( + serde_json::from_slice::(&populated).expect("parse populated page"), + json!({ + "meta": {"context_id": 1}, + "turns": [{"turn_id": 1}, {"turn_id": 2}], + "next_before_turn_id": 2 + }) + ); + } +} + +mod fast; + pub struct ProjectionResult { pub data: JsonValue, pub unknown: Option, } +pub fn project_turn( + item: &crate::store::TurnWithMeta, + registry: &Registry, + options: &TurnProjectionOptions<'_>, +) -> Result { + let declared_type_id = item.meta.declared_type_id.clone(); + let decoded_type = match options.type_hint_mode { + "explicit" => Ok(( + options + .as_type_id + .ok_or_else(|| StoreError::InvalidInput("as_type_id required".into()))? + .to_string(), + options + .as_type_version + .ok_or_else(|| StoreError::InvalidInput("as_type_version required".into()))?, + )), + "latest" => registry + .get_latest_type_version(&declared_type_id) + .map(|latest| (declared_type_id.clone(), latest.version)) + .ok_or_else(|| StoreError::NotFound("type descriptor".into())), + _ => Ok((declared_type_id.clone(), item.meta.declared_type_version)), + }; + let mut turn = Map::new(); + turn.insert( + "turn_id".into(), + render_id(item.record.turn_id, options.render.u64_format), + ); + turn.insert( + "parent_turn_id".into(), + render_id(item.record.parent_turn_id, options.render.u64_format), + ); + turn.insert("depth".into(), JsonValue::Number(item.record.depth.into())); + turn.insert("declared_type".into(), serde_json::json!({"type_id": declared_type_id, "type_version": item.meta.declared_type_version})); + if options.view == "typed" || options.view == "both" { + let projected = decoded_type.and_then(|(decoded_type_id, decoded_type_version)| { + let descriptor = registry + .get_type_version(&decoded_type_id, decoded_type_version) + .ok_or_else(|| StoreError::NotFound("type descriptor".into()))?; + let payload = item + .payload + .as_ref() + .ok_or_else(|| StoreError::InvalidInput("payload not loaded".into()))?; + let projected = project_msgpack(payload, descriptor, registry, options.render)?; + Ok((decoded_type_id, decoded_type_version, projected)) + }); + match projected { + Ok((decoded_type_id, decoded_type_version, projected)) => { + turn.insert( + "decoded_as".into(), + serde_json::json!({"type_id": decoded_type_id, "type_version": decoded_type_version}), + ); + turn.insert("data".into(), projected.data); + if let Some(unknown) = projected.unknown { + turn.insert("unknown".into(), unknown); + } + } + Err(error) => { + turn.insert( + "projection_error".into(), + serde_json::json!({"message": error.to_string()}), + ); + } + } + } + if options.view == "raw" || options.view == "both" { + let payload = item + .payload + .as_ref() + .ok_or_else(|| StoreError::InvalidInput("payload not loaded".into()))?; + turn.insert( + "content_hash_b3".into(), + JsonValue::String(hex::encode(item.record.payload_hash)), + ); + turn.insert( + "encoding".into(), + JsonValue::Number(item.meta.encoding.into()), + ); + turn.insert("compression".into(), JsonValue::Number(0u32.into())); + turn.insert( + "uncompressed_len".into(), + JsonValue::Number((payload.len() as u32).into()), + ); + match options.render.bytes_render { + BytesRender::Base64 => { + turn.insert( + "bytes_b64".into(), + JsonValue::String(base64::engine::general_purpose::STANDARD.encode(payload)), + ); + } + BytesRender::Hex => { + turn.insert("bytes_hex".into(), JsonValue::String(hex::encode(payload))); + } + BytesRender::LenOnly => { + turn.insert( + "bytes_len".into(), + JsonValue::Number((payload.len() as u64).into()), + ); + } + } + } + Ok(JsonValue::Object(turn)) +} + +fn render_id(value: u64, format: U64Format) -> JsonValue { + match format { + U64Format::String => JsonValue::String(value.to_string()), + U64Format::Number => JsonValue::Number(value.into()), + } +} + pub fn project_msgpack( payload: &[u8], descriptor: &TypeVersionSpec, @@ -60,24 +272,29 @@ pub fn project_msgpack( let mut cursor = std::io::Cursor::new(payload); let value = rmpv::decode::read_value(&mut cursor) .map_err(|e| StoreError::InvalidInput(format!("msgpack decode error: {e}")))?; + if cursor.position() != payload.len() as u64 { + return Err(StoreError::InvalidInput( + "payload contains trailing msgpack data".into(), + )); + } + if !matches!(value, Value::Map(_)) { + return Err(StoreError::InvalidInput("payload is not a map".into())); + } - let map = normalize_tags(&value)?; + let map = normalize_fields(&value, descriptor); let mut data = Map::new(); let mut unknown = Map::new(); for (tag, field) in descriptor.fields.iter() { - if let Some(val) = map.get(tag) { + if let Some(val) = map.known.get(tag) { let rendered = render_field_value(val, field, registry, options); data.insert(field.name.clone(), rendered); } } if options.include_unknown { - for (tag, val) in map.iter() { - if descriptor.fields.contains_key(tag) { - continue; - } - unknown.insert(tag.to_string(), render_value(val, options)); + for (tag, val) in map.unknown.iter() { + unknown.insert(tag.clone(), render_value(val, options)); } } @@ -91,20 +308,44 @@ pub fn project_msgpack( }) } -fn normalize_tags(value: &Value) -> Result> { - let mut out = HashMap::new(); +struct NormalizedFields { + known: HashMap, + unknown: HashMap, +} + +fn normalize_fields(value: &Value, descriptor: &TypeVersionSpec) -> NormalizedFields { + let mut known = HashMap::new(); + let mut unknown = HashMap::new(); let map = match value { Value::Map(map) => map, - _ => return Err(StoreError::InvalidInput("payload is not a map".into())), + _ => return NormalizedFields { known, unknown }, }; for (k, v) in map.iter() { - if let Some(tag) = key_to_tag(k) { - out.insert(tag, v.clone()); + let named = match k { + Value::String(name) => name.as_str().and_then(|name| { + descriptor + .fields + .iter() + .find_map(|(tag, field)| (field.name == name).then_some(*tag)) + }), + _ => None, + }; + if let Some(tag) = key_to_tag(k).or(named) { + if matches!(k, Value::Integer(_)) || !known.contains_key(&tag) { + known.insert(tag, v.clone()); + } + if !descriptor.fields.contains_key(&tag) { + let name = tag.to_string(); + if matches!(k, Value::Integer(_)) || !unknown.contains_key(&name) { + unknown.insert(name, v.clone()); + } + } + } else if let Value::String(name) = k { + unknown.insert(name.as_str().unwrap_or("").to_string(), v.clone()); } } - - Ok(out) + NormalizedFields { known, unknown } } fn key_to_tag(key: &Value) -> Option { @@ -157,7 +398,7 @@ fn render_field_value( match field_type { "u64" | "uint64" | "i64" | "int64" => render_u64(value, options), "u32" | "uint32" | "u8" | "uint8" | "int32" => render_int(value), - "string" => render_string(value), + "string" => render_string(value, options), "bool" => render_bool(value), "bytes" | "typed_blob" => render_bytes(value, options), "array" => render_array(value, field.items.as_ref(), registry, options), @@ -186,35 +427,29 @@ fn render_type_ref( return render_value(value, options); }; - // Normalize the value to a tag map - let Ok(map) = normalize_tags(value) else { + if !matches!(value, Value::Map(_)) { return render_value(value, options); - }; + } + + // Normalize the value to a tag map + let map = normalize_fields(value, type_spec); // Project using the type descriptor let mut data = Map::new(); for (tag, field) in type_spec.fields.iter() { - if let Some(val) = map.get(tag) { + if let Some(val) = map.known.get(tag) { let rendered = render_field_value(val, field, registry, options); data.insert(field.name.clone(), rendered); } } - // Propagate include_unknown into nested types — collect tags that the - // descriptor doesn't know about so they surface through the HTTP API. - if options.include_unknown { + if options.include_unknown && !map.unknown.is_empty() { let mut unknown = Map::new(); - for (tag, val) in map.iter() { - if type_spec.fields.contains_key(tag) { - continue; - } - unknown.insert(tag.to_string(), render_value(val, options)); - } - if !unknown.is_empty() { - data.insert("_unknown".into(), JsonValue::Object(unknown)); + for (name, value) in map.unknown { + unknown.insert(name, render_value(&value, options)); } + data.insert("_unknown".into(), JsonValue::Object(unknown)); } - JsonValue::Object(data) } @@ -233,7 +468,9 @@ fn render_value(value: &Value, options: &RenderOptions) -> JsonValue { } Value::F32(f) => JsonValue::Number(Number::from_f64(*f as f64).unwrap_or(Number::from(0))), Value::F64(f) => JsonValue::Number(Number::from_f64(*f).unwrap_or(Number::from(0))), - Value::String(s) => JsonValue::String(s.as_str().unwrap_or("").to_string()), + Value::String(s) => JsonValue::String( + limit_string(s.as_str().unwrap_or(""), options.string_limit).into_owned(), + ), Value::Binary(_) => render_bytes(value, options), Value::Array(arr) => { let items = arr.iter().map(|v| render_value(v, options)).collect(); @@ -258,13 +495,29 @@ fn render_value(value: &Value, options: &RenderOptions) -> JsonValue { } } -fn render_string(value: &Value) -> JsonValue { +fn render_string(value: &Value, options: &RenderOptions) -> JsonValue { match value { - Value::String(s) => JsonValue::String(s.as_str().unwrap_or("").to_string()), + Value::String(s) => JsonValue::String( + limit_string(s.as_str().unwrap_or(""), options.string_limit).into_owned(), + ), _ => JsonValue::Null, } } +pub fn limit_string(value: &str, limit: Option) -> Cow<'_, str> { + let Some(limit) = limit else { + return Cow::Borrowed(value); + }; + if value.len() <= limit { + return Cow::Borrowed(value); + } + let mut end = limit; + while end > 0 && !value.is_char_boundary(end) { + end -= 1; + } + Cow::Owned(value[..end].to_string()) +} + fn render_bool(value: &Value) -> JsonValue { match value { Value::Boolean(b) => JsonValue::Bool(*b), diff --git a/server/src/store.rs b/server/src/store.rs index 0dbec2f..06f230e 100644 --- a/server/src/store.rs +++ b/server/src/store.rs @@ -266,23 +266,65 @@ impl Store { include_payload: bool, ) -> Result> { let turns = self.turn_store.get_last(context_id, limit)?; - let mut out = Vec::with_capacity(turns.len()); - for record in turns { - let meta = self.turn_store.get_turn_meta(record.turn_id)?; - let payload = if include_payload { - Some(self.blob_store.get(&record.payload_hash)?) - } else { - None - }; - out.push(TurnWithMeta { - record, - meta, - payload, - }); + self.hydrate_turns(turns, include_payload) + } + + /// Hydrate turn metadata and payloads in page order. Blob reads are + /// coalesced by the CAS while duplicate payload hashes remain duplicated + /// in the returned page. + pub fn hydrate_turns( + &self, + turns: Vec, + include_payload: bool, + ) -> Result> { + let payloads: Vec>> = if include_payload { + let hashes: Vec<_> = turns.iter().map(|turn| turn.payload_hash).collect(); + self.blob_store + .get_many(&hashes)? + .into_iter() + .map(Some) + .collect() + } else { + std::iter::repeat_with(|| None).take(turns.len()).collect() + }; + turns + .into_iter() + .zip(payloads) + .map(|(record, payload)| { + let meta = self.turn_store.get_turn_meta(record.turn_id)?; + Ok(TurnWithMeta { + record, + meta, + payload, + }) + }) + .collect() + } + + /// Return one turn only when it belongs to the requested context. + pub fn get_turn(&self, turn_id: u64, include_payload: bool) -> Result { + let record = self.turn_store.get_turn(turn_id)?; + self.hydrate_turns(vec![record], include_payload)? + .into_iter() + .next() + .ok_or_else(|| StoreError::NotFound("turn".into())) + } + + /// Return an exact turn after validating context membership. + pub fn get_context_turn( + &self, + context_id: u64, + turn_id: u64, + include_payload: bool, + ) -> Result { + if !self.turn_store.context_contains_turn(context_id, turn_id)? { + return Err(StoreError::NotFound("turn in context".into())); } - Ok(out) + self.get_turn(turn_id, include_payload) } + /// Fetch a page before a cursor. Context membership is validated by the + /// turn store before the parent chain is traversed. pub fn get_before( &self, context_id: u64, @@ -293,21 +335,7 @@ impl Store { let turns = self .turn_store .get_before(context_id, before_turn_id, limit)?; - let mut out = Vec::with_capacity(turns.len()); - for record in turns { - let meta = self.turn_store.get_turn_meta(record.turn_id)?; - let payload = if include_payload { - Some(self.blob_store.get(&record.payload_hash)?) - } else { - None - }; - out.push(TurnWithMeta { - record, - meta, - payload, - }); - } - Ok(out) + self.hydrate_turns(turns, include_payload) } pub fn get_blob(&self, hash: &[u8; 32]) -> Result> { diff --git a/server/src/turn_store/mod.rs b/server/src/turn_store/mod.rs index 8a840be..48d13ed 100644 --- a/server/src/turn_store/mod.rs +++ b/server/src/turn_store/mod.rs @@ -12,6 +12,8 @@ use crc32fast::Hasher; use crate::error::{Result, StoreError}; +const ANCESTRY_BLOCK_SIZE: u32 = 256; + #[derive(Debug, Clone)] pub struct TurnRecord { pub turn_id: u64, @@ -57,6 +59,11 @@ pub struct TurnStore { turn_index: HashMap, turn_meta: HashMap, heads: HashMap, + /// Nearest block boundary ancestor for each turn. A map avoids allocating + /// a sparse vector when imported turn IDs are large. + ancestry_checkpoints: HashMap, + /// The head inherited when each context was created. + context_base_turns: HashMap, next_turn_id: u64, next_context_id: u64, @@ -108,11 +115,14 @@ impl TurnStore { turn_index: HashMap::new(), turn_meta: HashMap::new(), heads: HashMap::new(), + ancestry_checkpoints: HashMap::new(), + context_base_turns: HashMap::new(), next_turn_id: 1, next_context_id: 1, }; store.load_turns()?; + store.rebuild_ancestry_checkpoints()?; store.load_meta()?; store.load_heads()?; store.rebuild_index()?; @@ -163,6 +173,38 @@ impl TurnStore { Ok(()) } + fn rebuild_ancestry_checkpoints(&mut self) -> Result<()> { + self.ancestry_checkpoints.clear(); + let mut turn_ids: Vec<_> = self.turns.keys().copied().collect(); + turn_ids.sort_unstable(); + for turn_id in turn_ids { + let record = self + .turns + .get(&turn_id) + .ok_or_else(|| StoreError::Corrupt("turn disappeared during index build".into()))?; + let checkpoint = if record.depth % ANCESTRY_BLOCK_SIZE == 0 { + record.turn_id + } else if record.parent_turn_id == 0 { + return Err(StoreError::Corrupt( + "turn ancestry checkpoint is incomplete".into(), + )); + } else { + *self + .ancestry_checkpoints + .get(&record.parent_turn_id) + .ok_or_else(|| { + StoreError::Corrupt("turn ancestry checkpoint is incomplete".into()) + })? + }; + self.ancestry_checkpoints.insert(turn_id, checkpoint); + } + Ok(()) + } + + fn checkpoint(&self, turn_id: u64) -> Option { + self.ancestry_checkpoints.get(&turn_id).copied() + } + fn load_meta(&mut self) -> Result<()> { self.turn_meta.clear(); self.turns_meta.seek(SeekFrom::Start(0))?; @@ -234,6 +276,7 @@ impl TurnStore { fn load_heads(&mut self) -> Result<()> { self.heads.clear(); + self.context_base_turns.clear(); self.heads_tbl.seek(SeekFrom::Start(0))?; loop { let start = self.heads_tbl.stream_position()?; @@ -302,6 +345,9 @@ impl TurnStore { flags, }, ); + self.context_base_turns + .entry(context_id) + .or_insert(head_turn_id); } Ok(()) } @@ -357,6 +403,7 @@ impl TurnStore { self.write_head(&head)?; self.heads.insert(context_id, head.clone()); + self.context_base_turns.insert(context_id, head_turn_id); Ok(head) } @@ -454,6 +501,14 @@ impl TurnStore { ); self.turns.insert(turn_id, record.clone()); self.turn_index.insert(turn_id, offset); + let checkpoint = if depth % ANCESTRY_BLOCK_SIZE == 0 { + turn_id + } else { + self.checkpoint(parent_id).ok_or_else(|| { + StoreError::Corrupt("parent ancestry checkpoint is missing".into()) + })? + }; + self.ancestry_checkpoints.insert(turn_id, checkpoint); // update head let head = ContextHead { @@ -493,6 +548,44 @@ impl TurnStore { .ok_or_else(|| StoreError::NotFound("turn".into())) } + /// Check membership without walking from the head one edge at a time for + /// every block of a deep ancestry chain. + pub fn context_contains_turn(&self, context_id: u64, turn_id: u64) -> Result { + let head = self + .heads + .get(&context_id) + .ok_or_else(|| StoreError::NotFound("context".into()))?; + let target = self + .turns + .get(&turn_id) + .ok_or_else(|| StoreError::NotFound("turn".into()))?; + if target.depth > head.head_depth { + return Ok(false); + } + let mut current_id = head.head_turn_id; + let mut current_depth = head.head_depth; + while current_depth > target.depth { + let distance = current_depth - target.depth; + let within_block = current_depth % ANCESTRY_BLOCK_SIZE; + if within_block > 0 && within_block <= distance { + if let Some(checkpoint) = self.checkpoint(current_id) { + if checkpoint != current_id { + current_id = checkpoint; + current_depth -= within_block; + continue; + } + } + } + let current = self + .turns + .get(¤t_id) + .ok_or_else(|| StoreError::Corrupt("context ancestry is incomplete".into()))?; + current_id = current.parent_turn_id; + current_depth -= 1; + } + Ok(current_id == turn_id) + } + pub fn get_turn_meta(&self, turn_id: u64) -> Result { self.turn_meta .get(&turn_id) @@ -532,7 +625,14 @@ impl TurnStore { .get(&context_id) .ok_or_else(|| StoreError::NotFound("context".into()))?; - if before_turn_id == 0 || head.head_turn_id == 0 { + if before_turn_id == 0 { + return self.get_last(context_id, limit); + } + + if !self.context_contains_turn(context_id, before_turn_id)? { + return Err(StoreError::NotFound("turn in context".into())); + } + if head.head_turn_id == 0 { return self.get_last(context_id, limit); } @@ -562,14 +662,33 @@ impl TurnStore { .get(&context_id) .ok_or_else(|| StoreError::NotFound("context".into()))?; - // Walk back from head to find the turn with depth=0 + let base_turn_id = self + .context_base_turns + .get(&context_id) + .copied() + .ok_or_else(|| StoreError::NotFound("context base".into()))?; + if head.head_turn_id == base_turn_id { + return Err(StoreError::NotFound("first turn".into())); + } + + // Walk back to the first turn owned by this context. let mut current = head.head_turn_id; while current != 0 { let rec = self .turns .get(¤t) .ok_or_else(|| StoreError::NotFound("turn".into()))?; - if rec.depth == 0 { + if base_turn_id == 0 && rec.depth == 0 { + return Ok(rec.clone()); + } + if rec.turn_id == base_turn_id { + return Err(StoreError::NotFound("first turn".into())); + } + let parent = self.turns.get(&rec.parent_turn_id); + if parent + .map(|record| record.turn_id == base_turn_id) + .unwrap_or(false) + { return Ok(rec.clone()); } current = rec.parent_turn_id; @@ -581,7 +700,7 @@ impl TurnStore { pub fn list_recent_contexts(&self, limit: u32) -> Vec { let mut contexts: Vec = self.heads.values().cloned().collect(); // Sort by created_at descending (most recent first) - contexts.sort_by(|a, b| b.created_at_unix_ms.cmp(&a.created_at_unix_ms)); + contexts.sort_by_key(|context| std::cmp::Reverse(context.created_at_unix_ms)); contexts.truncate(limit as usize); contexts } diff --git a/server/tests/registry_projection.rs b/server/tests/registry_projection.rs index 09843cb..b7779ea 100644 --- a/server/tests/registry_projection.rs +++ b/server/tests/registry_projection.rs @@ -1,9 +1,11 @@ // Copyright 2025 StrongDM Inc // SPDX-License-Identifier: Apache-2.0 -use cxdb_server::projection::project_msgpack; +use cxdb_server::projection::{project_msgpack, project_turn_page, serialize_turn_page}; use cxdb_server::projection::{BytesRender, EnumRender, RenderOptions, TimeRender, U64Format}; use cxdb_server::registry::Registry; +use cxdb_server::store::TurnWithMeta; +use cxdb_server::turn_store::{TurnMeta, TurnRecord}; use rmpv::Value; use tempfile::tempdir; @@ -14,6 +16,7 @@ fn default_options() -> RenderOptions { enum_render: EnumRender::Label, time_render: TimeRender::Iso, include_unknown: true, + string_limit: None, } } @@ -68,6 +71,7 @@ fn registry_ingest_and_project() { enum_render: EnumRender::Label, time_render: TimeRender::Iso, include_unknown: true, + string_limit: None, }; let projection = project_msgpack(&buf, desc, ®istry, &options).expect("project"); @@ -733,3 +737,167 @@ fn array_ref_items_include_unknown_tags() { "_unknown should not appear when there are no unknown tags" ); } + +#[test] +fn named_keys_are_recursive_numeric_priority_and_utf8_bounded() { + let dir = tempdir().expect("tempdir"); + let mut registry = Registry::open(dir.path()).expect("open registry"); + let bundle = r#"{ + "registry_version": 1, "bundle_id": "named#test", + "types": {"test:Message": {"versions": {"1": {"fields": { + "1": {"name": "text", "type": "string"}, + "2": {"name": "count", "type": "u64"} + }}}}} + }"#; + registry + .put_bundle("named#test", bundle.as_bytes()) + .expect("bundle"); + let desc = registry.get_type_version("test:Message", 1).expect("desc"); + let value = Value::Map(vec![ + (Value::String("text".into()), Value::String("named".into())), + (Value::Integer(1.into()), Value::String("numeric".into())), + (Value::String("1".into()), Value::String("alias".into())), + ( + Value::String("extra".into()), + Value::String("éclair".into()), + ), + (Value::String("count".into()), Value::Integer(7.into())), + ]); + let mut payload = Vec::new(); + rmpv::encode::write_value(&mut payload, &value).expect("encode"); + + let bounded = RenderOptions { + string_limit: Some(4), + ..default_options() + }; + let projected = project_msgpack(&payload, desc, ®istry, &bounded).expect("project"); + assert_eq!(projected.data["text"], "nume"); + assert_eq!(projected.data["count"], "7"); + assert_eq!(projected.unknown.as_ref().unwrap()["extra"], "écl"); + + let record = TurnRecord { + turn_id: 9, + parent_turn_id: 0, + depth: 0, + codec: 1, + type_tag: 0, + payload_hash: *blake3::hash(&payload).as_bytes(), + flags: 0, + created_at_unix_ms: 0, + }; + let item = TurnWithMeta { + record, + meta: TurnMeta { + declared_type_id: "test:Message".into(), + declared_type_version: 1, + encoding: 1, + compression: 0, + uncompressed_len: payload.len() as u32, + }, + payload: Some(payload), + }; + let numeric = RenderOptions { + u64_format: U64Format::Number, + ..bounded.clone() + }; + let options = cxdb_server::projection::TurnProjectionOptions { + view: "typed", + type_hint_mode: "inherit", + as_type_id: None, + as_type_version: None, + render: &numeric, + }; + let normal = + project_turn_page(std::slice::from_ref(&item), ®istry, &options).expect("normal"); + let fast = serialize_turn_page(std::slice::from_ref(&item), ®istry, &options).expect("fast"); + let fast_json: serde_json::Value = serde_json::from_slice(&fast[0]).expect("fast json"); + assert_eq!(fast_json, normal[0]); + assert_eq!(fast_json["turn_id"], 9); +} + +#[test] +fn non_map_type_ref_falls_back_to_raw_value() { + let dir = tempdir().expect("tempdir"); + let mut registry = Registry::open(dir.path()).expect("open registry"); + let bundle = r#"{ + "registry_version": 1, + "bundle_id": "non-map-ref", + "types": { + "test:Outer": {"versions": {"1": {"fields": { + "1": {"name": "nested", "type": "ref", "ref": "test:Inner"} + }}}}, + "test:Inner": {"versions": {"1": {"fields": { + "1": {"name": "value", "type": "string"} + }}}} + } + }"#; + registry + .put_bundle("non-map-ref", bundle.as_bytes()) + .expect("bundle"); + let descriptor = registry + .get_type_version("test:Outer", 1) + .expect("descriptor"); + let value = Value::Map(vec![( + Value::Integer(1.into()), + Value::String("legacy scalar".into()), + )]); + let mut payload = Vec::new(); + rmpv::encode::write_value(&mut payload, &value).expect("encode"); + + let projected = + project_msgpack(&payload, descriptor, ®istry, &default_options()).expect("project"); + assert_eq!(projected.data["nested"], "legacy scalar"); +} + +#[test] +fn invalid_turn_payload_is_isolated_in_typed_pages() { + let dir = tempdir().expect("tempdir"); + let mut registry = Registry::open(dir.path()).expect("open registry"); + let bundle = r#"{ + "registry_version": 1, + "bundle_id": "isolation", + "types": {"test:Message": {"versions": {"1": {"fields": { + "1": {"name": "text", "type": "string"} + }}}}} + }"#; + registry + .put_bundle("isolation", bundle.as_bytes()) + .expect("bundle"); + let payload = b"not-msgpack".to_vec(); + let item = TurnWithMeta { + record: TurnRecord { + turn_id: 10, + parent_turn_id: 9, + depth: 2, + codec: 1, + type_tag: 0, + payload_hash: *blake3::hash(&payload).as_bytes(), + flags: 0, + created_at_unix_ms: 0, + }, + meta: TurnMeta { + declared_type_id: "test:Message".into(), + declared_type_version: 1, + encoding: 1, + compression: 0, + uncompressed_len: payload.len() as u32, + }, + payload: Some(payload), + }; + let render = default_options(); + let options = cxdb_server::projection::TurnProjectionOptions { + view: "typed", + type_hint_mode: "inherit", + as_type_id: None, + as_type_version: None, + render: &render, + }; + + let projected = + project_turn_page(std::slice::from_ref(&item), ®istry, &options).expect("isolated page"); + assert_eq!(projected[0]["turn_id"], "10"); + assert!(projected[0]["projection_error"]["message"].is_string()); + let serialized = serialize_turn_page(&[item], ®istry, &options).expect("serialized page"); + let serialized: serde_json::Value = serde_json::from_slice(&serialized[0]).expect("turn json"); + assert!(serialized["projection_error"]["message"].is_string()); +} diff --git a/server/tests/store_basic.rs b/server/tests/store_basic.rs index 40e44c9..f1847a3 100644 --- a/server/tests/store_basic.rs +++ b/server/tests/store_basic.rs @@ -157,6 +157,132 @@ fn indexes_parent_child_context_lineage() { assert_eq!(descendants, vec![grandchild.context_id, child.context_id]); } +#[test] +fn exact_turn_and_cursor_are_context_scoped() { + let dir = tempdir().expect("tempdir"); + let mut store = Store::open(dir.path()).expect("open store"); + let root = store.create_context(0).expect("root"); + let payload = b"root"; + let hash = blake3::hash(payload); + let (root_turn, _) = store + .append_turn( + root.context_id, + 0, + "test:Message".into(), + 1, + 1, + 0, + payload.len() as u32, + *hash.as_bytes(), + payload, + ) + .expect("root turn"); + let child = store.fork_context(root_turn.turn_id).expect("fork"); + let child_payload = b"child"; + let child_hash = blake3::hash(child_payload); + let (child_turn, _) = store + .append_turn( + child.context_id, + 0, + "test:Message".into(), + 1, + 1, + 0, + child_payload.len() as u32, + *child_hash.as_bytes(), + child_payload, + ) + .expect("child turn"); + + assert!(store + .get_context_turn(child.context_id, child_turn.turn_id, true) + .is_ok()); + assert!(store + .get_context_turn(root.context_id, child_turn.turn_id, true) + .is_err()); + assert!(store + .get_before(root.context_id, child_turn.turn_id, 10, true) + .is_err()); + assert_eq!( + store + .get_context_turn(child.context_id, root_turn.turn_id, true) + .expect("inherited turn") + .payload + .as_deref(), + Some(payload.as_slice()) + ); +} + +#[test] +fn deep_ancestry_membership_uses_checkpoints() { + let dir = tempdir().expect("tempdir"); + let mut store = Store::open(dir.path()).expect("open store"); + let context = store.create_context(0).expect("context"); + let payload = b"same"; + let hash = blake3::hash(payload); + let mut ids = Vec::new(); + for _ in 0..600 { + let (turn, _) = store + .append_turn( + context.context_id, + 0, + "test:Message".into(), + 1, + 1, + 0, + payload.len() as u32, + *hash.as_bytes(), + payload, + ) + .expect("append"); + ids.push(turn.turn_id); + } + assert!(store + .get_context_turn(context.context_id, ids[0], false) + .is_ok()); + assert!(store + .get_context_turn(context.context_id, ids[599], false) + .is_ok()); +} + +#[test] +fn sequential_append_failure_keeps_prior_success() { + let dir = tempdir().expect("tempdir"); + let mut store = Store::open(dir.path()).expect("open store"); + let context = store.create_context(0).expect("context"); + let payload = b"first"; + let hash = blake3::hash(payload); + store + .append_turn( + context.context_id, + 0, + "test:Message".into(), + 1, + 1, + 0, + payload.len() as u32, + *hash.as_bytes(), + payload, + ) + .expect("first append"); + let failed = store.append_turn( + context.context_id, + 0, + "test:Message".into(), + 1, + 1, + 0, + payload.len() as u32, + [0; 32], + payload, + ); + assert!(failed.is_err()); + assert_eq!( + store.get_last(context.context_id, 10, true).unwrap().len(), + 1 + ); +} + fn encode_context_metadata_payload( parent_context_id: Option, root_context_id: Option,