Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 23 additions & 0 deletions .github/workflows/ci.yml
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,18 @@ jobs:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- name: Set cache key date
id: cache-date
run: echo "date=$(date -u +'%Y-%m-%d')" >> $GITHUB_OUTPUT
- name: Cache cargo-audit binary and advisory DB
uses: actions/cache@v4
with:
path: |
~/.cargo/bin/cargo-audit
~/.cargo/advisory-db
key: ${{ runner.os }}-cargo-audit-${{ steps.cache-date.outputs.date }}
restore-keys: |
${{ runner.os }}-cargo-audit-
# A prebuilt binary, deliberately not rustsec/audit-check: that action shells out to
# `cargo install cargo-audit`, which picks up the 1.84.1 pin in rust-toolchain.toml and
# fails, because a current cargo-audit needs rustc 1.88+. Installing a prebuilt binary
Expand All @@ -98,6 +110,17 @@ jobs:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- name: Set cache key date
id: cache-date
run: echo "date=$(date -u +'%Y-%m-%d')" >> $GITHUB_OUTPUT
- name: Cache cargo-deny advisory DB
uses: actions/cache@v4
with:
path: |
~/.cargo/advisory-dbs
key: ${{ runner.os }}-cargo-deny-${{ steps.cache-date.outputs.date }}
restore-keys: |
${{ runner.os }}-cargo-deny-
- uses: EmbarkStudios/cargo-deny-action@v2

secret-scan:
Expand Down
1 change: 1 addition & 0 deletions crates/api/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@ url = "2"
[dev-dependencies]
tokio.workspace = true
tower.workspace = true
tracing-subscriber.workspace = true
sqlx.workspace = true
stellar-base.workspace = true
dotenvy = "0.15"
Expand Down
40 changes: 40 additions & 0 deletions crates/api/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -19,13 +19,52 @@ pub use error::{ApiError, ApiResult, Envelope};
pub use state::AppState;

use axum::extract::{DefaultBodyLimit, Request, State};
use axum::http::header::HeaderName;
use axum::http::StatusCode;
use axum::http::HeaderValue;
use axum::middleware::{self, Next};
use axum::response::{IntoResponse, Response};
use axum::routing::{delete, get, patch, post};
use axum::{Json, Router};
use std::time::Duration;
use tower_http::cors::{Any, CorsLayer};
use tracing::Instrument;

/// Canonical header name for request correlation IDs.
pub static REQUEST_ID_HEADER: HeaderName = HeaderName::from_static("x-request-id");

/// Extracted or generated request ID stored in request extensions.
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RequestId(pub String);

/// Middleware that extracts or generates a request ID, enters an info span, and echoes it in responses.
pub async fn request_id_middleware(mut req: Request, next: Next) -> Response {
// Reuse caller-supplied X-Request-Id if non-empty and valid ASCII, else generate UUIDv4.
let request_id = req
.headers()
.get(&REQUEST_ID_HEADER)
.and_then(|v| v.to_str().ok())
.map(str::trim)
.filter(|s| !s.is_empty() && HeaderValue::from_str(s).is_ok())
.map(|s| s.to_string())
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());

// Record request ID in request extensions for handlers.
req.extensions_mut().insert(RequestId(request_id.clone()));

// Open tracing span carrying the request ID for cross-service log correlation.
let span = tracing::info_span!("request", request_id = %request_id);

// Run downstream middleware and routes within the span context.
let mut response = next.run(req).instrument(span).await;

// Attach request ID header to response.
if let Ok(val) = HeaderValue::from_str(&request_id) {
response.headers_mut().insert(REQUEST_ID_HEADER.clone(), val);
}

response
}

/// Keep API request payloads bounded to a deliberate, documented ceiling.
///
Expand Down Expand Up @@ -218,6 +257,7 @@ pub fn build_router(state: AppState) -> Router {
// cleanly with `Router::layer` here.
.layer(DefaultBodyLimit::max(REQUEST_BODY_LIMIT))
.layer(cors)
.layer(middleware::from_fn(request_id_middleware))
.with_state(state)
}

Expand Down
166 changes: 166 additions & 0 deletions crates/api/tests/request_id_tests.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,166 @@
//! Tests for per-request id propagation through middleware, response headers, and tracing spans.

use axum::body::Body;
use axum::http::header::HeaderName;
use axum::http::{Request, StatusCode};
use axum::routing::get;
use axum::Router;
use octo_api::{build_router, request_id_middleware, AppState, REQUEST_ID_HEADER};
use octo_store::Store;
use octo_wallet_core::StellarNetwork;
use std::sync::{Arc, Mutex, Once};
use tower::ServiceExt;
use tracing_subscriber::fmt::MakeWriter;
use uuid::Uuid;

static LOAD_ENV: Once = Once::new();

fn database_url() -> Option<String> {
LOAD_ENV.call_once(|| {
let _ = dotenvy::dotenv();
});
std::env::var("DATABASE_URL").ok()
}

async fn test_state() -> Option<AppState> {
let url = database_url()?;
let store = Store::connect(&url).await.ok()?;
let _ = store.migrate().await;
let master_key = [42u8; 32];
Some(AppState::new(
store,
master_key,
StellarNetwork::Testnet,
"https://horizon-testnet.stellar.org".into(),
None,
octo_email::EmailSender::new_captured(),
))
}

async fn test_router() -> Router {
if let Some(state) = test_state().await {
build_router(state)
} else {
// Fallback router with identical middleware layer when database is unavailable.
Router::new()
.route("/health", get(|| async { "ok" }))
.layer(axum::middleware::from_fn(request_id_middleware))
}
}

#[derive(Clone)]
struct BufferWriter(Arc<Mutex<Vec<u8>>>);

impl std::io::Write for BufferWriter {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.0.lock().unwrap().extend_from_slice(buf);
Ok(buf.len())
}

fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}

impl<'a> MakeWriter<'a> for BufferWriter {
type Writer = BufferWriter;

fn make_writer(&'a self) -> Self::Writer {
self.clone()
}
}

#[tokio::test]
async fn every_response_carries_an_x_request_id_header() {
let app = test_router().await;
let req = Request::builder()
.uri("/health")
.body(Body::empty())
.unwrap();

let resp = app.oneshot(req).await.expect("execute request");
assert_eq!(resp.status(), StatusCode::OK);

// Verify response carries X-Request-Id header.
let header_val = resp
.headers()
.get(&REQUEST_ID_HEADER)
.expect("X-Request-Id header present in response")
.to_str()
.expect("header is valid ASCII");

// Verify generated request id is a valid UUIDv4.
let parsed = Uuid::parse_str(header_val);
assert!(parsed.is_ok(), "generated request id must be a valid UUID");
assert_eq!(parsed.unwrap().get_version_num(), 4);
}

#[tokio::test]
async fn a_caller_supplied_x_request_id_is_echoed_back_unchanged() {
let app = test_router().await;
let custom_id = "client-trace-777-custom-id";
let req = Request::builder()
.uri("/health")
.header(HeaderName::from_static("x-request-id"), custom_id)
.body(Body::empty())
.unwrap();

let resp = app.oneshot(req).await.expect("execute request");
assert_eq!(resp.status(), StatusCode::OK);

// Verify caller-supplied request ID is preserved exactly.
let header_val = resp
.headers()
.get(&REQUEST_ID_HEADER)
.expect("X-Request-Id header present in response")
.to_str()
.expect("header is valid ASCII");
assert_eq!(header_val, custom_id);
}

#[tokio::test]
async fn log_output_for_a_request_consistently_carries_the_same_request_id_across_nested_spans() {
let log_buffer = Arc::new(Mutex::new(Vec::new()));
let subscriber = tracing_subscriber::fmt()
.with_writer(BufferWriter(log_buffer.clone()))
.with_ansi(false)
.finish();

// Register test subscriber for the current thread during test execution.
let _guard = tracing::subscriber::set_default(subscriber);

// Handler with a nested child span to test span inheritance.
async fn nested_handler() -> &'static str {
let child_span = tracing::info_span!("horizon_dispatch", operation = "poll_status");
let _enter = child_span.enter();
tracing::info!("nested horizon call completed");
"ok"
}

let app = Router::new()
.route("/trace-test", get(nested_handler))
.layer(axum::middleware::from_fn(request_id_middleware));

let custom_id = "trace-correlation-id-9988";
let req = Request::builder()
.uri("/trace-test")
.header(HeaderName::from_static("x-request-id"), custom_id)
.body(Body::empty())
.unwrap();

let resp = app.oneshot(req).await.expect("execute request");
assert_eq!(resp.status(), StatusCode::OK);

// Extract captured logs and assert request id is threaded through nested spans.
let logs = String::from_utf8(log_buffer.lock().unwrap().clone()).expect("valid utf8 logs");
assert!(
logs.contains(custom_id),
"log output must contain the request id: {}",
logs
);
assert!(
logs.contains("nested horizon call completed"),
"log output must contain nested span event: {}",
logs
);
}
Loading
Loading