diff --git a/Cargo.lock b/Cargo.lock index 0f21dca81..9ce9c23f9 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -492,12 +492,56 @@ dependencies = [ "libc", ] +[[package]] +name = "anstream" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "824a212faf96e9acacdbd09febd34438f8f711fb84e09a8916013cd7815ca28d" +dependencies = [ + "anstyle", + "anstyle-parse", + "anstyle-query", + "anstyle-wincon", + "colorchoice", + "is_terminal_polyfill", + "utf8parse", +] + [[package]] name = "anstyle" version = "1.0.13" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5192cca8006f1fd4f7237516f40fa183bb07f8fbdfedaa0036de5ea9b0b45e78" +[[package]] +name = "anstyle-parse" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "52ce7f38b242319f7cabaa6813055467063ecdc9d355bbb4ce0c68908cd8130e" +dependencies = [ + "utf8parse", +] + +[[package]] +name = "anstyle-query" +version = "1.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc" +dependencies = [ + "windows-sys 0.61.2", +] + +[[package]] +name = "anstyle-wincon" +version = "3.0.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "291e6a250ff86cd4a820112fb8898808a366d8f9f58ce16d1f538353ad55747d" +dependencies = [ + "anstyle", + "once_cell_polyfill", + "windows-sys 0.61.2", +] + [[package]] name = "anthropic_compat" version = "0.0.0" @@ -527,6 +571,7 @@ dependencies = [ "bhttp", "bytes", "chrono", + "clap", "config", "database", "deadpool-postgres", @@ -1971,6 +2016,46 @@ dependencies = [ "zeroize", ] +[[package]] +name = "clap" +version = "4.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ddb117e43bbf7dacf0a4190fef4d345b9bad68dfc649cb349e7d17d28428e51" +dependencies = [ + "clap_builder", + "clap_derive", +] + +[[package]] +name = "clap_builder" +version = "4.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "714a53001bf66416adb0e2ef5ac857140e7dc3a0c48fb28b2f10762fc4b5069f" +dependencies = [ + "anstream", + "anstyle", + "clap_lex", + "strsim", +] + +[[package]] +name = "clap_derive" +version = "4.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f2ce8604710f6733aa641a2b3731eaa1e8b3d9973d5e3565da11800813f997a9" +dependencies = [ + "heck", + "proc-macro2", + "quote", + "syn 2.0.117", +] + +[[package]] +name = "clap_lex" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c8d4a3bb8b1e0c1050499d1815f5ab16d04f0959b233085fb31653fbfc9d98f9" + [[package]] name = "cmake" version = "0.1.57" @@ -1986,6 +2071,12 @@ version = "0.5.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3f88a43d011fc4a6876cb7344703e297c71dda42494fee094d5f7c76bf13f746" +[[package]] +name = "colorchoice" +version = "1.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570" + [[package]] name = "combine" version = "4.6.7" @@ -2383,8 +2474,10 @@ checksum = "2a2330da5de22e8a3cb63252ce2abb30116bf5265e89c0e01bc17015ce30a476" name = "database" version = "0.0.0" dependencies = [ + "aes-gcm", "anyhow", "async-trait", + "base64 0.22.1", "chrono", "config", "deadpool", @@ -3684,7 +3777,7 @@ dependencies = [ "libc", "percent-encoding", "pin-project-lite", - "socket2 0.5.10", + "socket2 0.6.4", "system-configuration", "tokio", "tower-service", @@ -3954,6 +4047,12 @@ dependencies = [ "serde", ] +[[package]] +name = "is_terminal_polyfill" +version = "1.70.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6cb138bb79a146c1bd460005623e142ef0181e3d0219cb493e02f7d08a35695" + [[package]] name = "itertools" version = "0.10.5" @@ -4750,6 +4849,12 @@ dependencies = [ "portable-atomic", ] +[[package]] +name = "once_cell_polyfill" +version = "1.70.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" + [[package]] name = "opaque-debug" version = "0.3.1" @@ -5432,7 +5537,7 @@ dependencies = [ "quinn-udp", "rustc-hash", "rustls", - "socket2 0.5.10", + "socket2 0.6.4", "thiserror 2.0.18", "tokio", "tracing", @@ -5470,7 +5575,7 @@ dependencies = [ "cfg_aliases", "libc", "once_cell", - "socket2 0.5.10", + "socket2 0.6.4", "tracing", "windows-sys 0.60.2", ] @@ -7658,6 +7763,12 @@ version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" +[[package]] +name = "utf8parse" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" + [[package]] name = "utoipa" version = "5.5.0" @@ -7966,7 +8077,7 @@ version = "0.1.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" dependencies = [ - "windows-sys 0.48.0", + "windows-sys 0.61.2", ] [[package]] diff --git a/Dockerfile b/Dockerfile index cdcd4d847..98ef07c1c 100644 --- a/Dockerfile +++ b/Dockerfile @@ -40,7 +40,7 @@ COPY crates/ ./crates/ COPY .cargo/ ./.cargo/ # Build the application in release mode -RUN cargo build --release --locked --bin api +RUN cargo build --release --locked --bin api --bin database_encryption_worker # Runtime stage @@ -86,6 +86,7 @@ WORKDIR /app # Copy the built binary COPY --from=builder /app/target/release/api /app/api +COPY --from=builder /app/target/release/database_encryption_worker /app/database_encryption_worker # Copy the migration SQL files RUN mkdir -p /app/crates/database/src/migrations/sql diff --git a/crates/api/Cargo.toml b/crates/api/Cargo.toml index be184073d..90e7c8058 100644 --- a/crates/api/Cargo.toml +++ b/crates/api/Cargo.toml @@ -16,7 +16,9 @@ tokio-stream = "0.1" tracing = "0.1" tracing-subscriber = { version = "0.3", features = ["env-filter", "json"] } tokio = { version = "1", features = ["full"] } +tokio-postgres = "0.7.18" anyhow = "1.0" +clap = { version = "4.5", features = ["derive"] } config = { path = "../config" } database = { path = "../database" } services = { path = "../services" } @@ -73,7 +75,6 @@ dotenvy = "0.15.7" k256 = { version = "0.13", features = ["ecdsa", "arithmetic"] } sha3 = "0.12" hmac = "0.13" -tokio-postgres = "0.7.18" deadpool-postgres = { version = "0.14", features = ["rt_tokio_1"] } ed25519-dalek = { version = "2.1", features = ["rand_core"] } rand = "0.10" diff --git a/crates/api/src/bin/database_encryption_worker.rs b/crates/api/src/bin/database_encryption_worker.rs new file mode 100644 index 000000000..189ba1d93 --- /dev/null +++ b/crates/api/src/bin/database_encryption_worker.rs @@ -0,0 +1,107 @@ +use anyhow::{Context, Result}; +use api::database_encryption::{ + operational_migrate, operational_scan, operational_verify, DatabaseEncryptionState, +}; +use clap::{Parser, Subcommand}; +use database::Database; +use uuid::Uuid; + +#[derive(Parser)] +#[command(about = "One-off database encryption backfill worker")] +struct Cli { + #[command(subcommand)] + command: Command, +} + +#[derive(Subcommand)] +enum Command { + Scan { + #[arg(long, value_delimiter = ',')] + scope: Vec, + }, + Migrate { + #[arg(long, required = true, value_delimiter = ',')] + scope: Vec, + #[arg(long, default_value_t = 500)] + batch_size: i64, + #[arg(long)] + max_rows: Option, + #[arg(long)] + resume: Option, + #[arg(long)] + operator: String, + }, + Verify { + #[arg(long, value_delimiter = ',')] + scope: Vec, + }, +} + +#[tokio::main] +async fn main() { + if run().await.is_err() { + println!( + "DATABASE_ENCRYPTION_WORKER_RESULT {}", + serde_json::json!({"status":"failed","error_class":"worker_failed"}) + ); + std::process::exit(1); + } +} + +async fn run() -> Result<()> { + let cli = Cli::parse(); + let database_config = config::DatabaseConfig::from_env() + .map_err(anyhow::Error::msg) + .context("invalid database configuration")?; + let key = read_encryption_key()?; + let database = Database::from_config(&database_config).await?; + let state = DatabaseEncryptionState::new(database.pool().clone(), &key)?; + + match cli.command { + Command::Scan { scope } => { + println!( + "{}", + serde_json::to_string_pretty(&operational_scan(&state, scope).await?)? + ); + print_success(None); + } + Command::Verify { scope } => { + let report = operational_verify(&state, scope).await?; + println!("{}", serde_json::to_string_pretty(&report)?); + if report["pass"] != true { + anyhow::bail!("verification found plaintext or invalid envelopes"); + } + print_success(None); + } + Command::Migrate { + scope, + batch_size, + max_rows, + resume, + operator, + } => { + let id = + operational_migrate(&state, scope, batch_size, max_rows, resume, &operator).await?; + print_success(Some(id)); + } + } + Ok(()) +} + +fn print_success(job_id: Option) { + println!( + "DATABASE_ENCRYPTION_WORKER_RESULT {}", + serde_json::json!({"status":"completed","job_id":job_id}) + ); +} + +fn read_encryption_key() -> Result { + if let Ok(key) = std::env::var("S3_ENCRYPTION_KEY") { + return Ok(key); + } + let path = std::env::var("S3_ENCRYPTION_KEY_FILE") + .context("S3_ENCRYPTION_KEY or S3_ENCRYPTION_KEY_FILE is required")?; + std::fs::read_to_string(path) + .context("failed to read S3_ENCRYPTION_KEY_FILE") + .map(|key| key.trim().to_string()) +} diff --git a/crates/api/src/database_encryption.rs b/crates/api/src/database_encryption.rs new file mode 100644 index 000000000..f85eb90c3 --- /dev/null +++ b/crates/api/src/database_encryption.rs @@ -0,0 +1,953 @@ +use crate::database_encryption_inventory::{policy_for, Classification}; +use crate::middleware::AdminUser; +use axum::{extract::Path, http::StatusCode, Extension, Json}; +use chrono::{DateTime, Utc}; +use database::DbPool; +use serde::{Deserialize, Serialize}; +use serde_json::{json, Value}; +use uuid::Uuid; + +const MARKER: &str = "__near_db_encrypted"; + +// Confidential fields and API contracts. + +#[derive(Clone, Copy, Debug, Serialize)] +#[serde(rename_all = "snake_case")] +enum Kind { + Text, + Json, +} + +#[derive(Clone, Copy, Debug, Serialize)] +struct Field { + table: &'static str, + column: &'static str, + kind: Kind, + action: &'static str, + reason: &'static str, +} + +const fn f(table: &'static str, column: &'static str, kind: Kind, reason: &'static str) -> Field { + Field { + table, + column, + kind, + action: "encrypt", + reason, + } +} + +const FIELDS: &[Field] = &[ + f( + "response_items", + "item", + Kind::Json, + "Persisted transcript and tool payloads", + ), + f( + "responses", + "instructions", + Kind::Text, + "Prompt instructions", + ), + f( + "responses", + "metadata", + Kind::Json, + "User response metadata", + ), + f( + "conversations", + "metadata", + Kind::Json, + "Conversation title and metadata", + ), + f("files", "filename", Kind::Text, "User file metadata"), + f("files", "storage_key", Kind::Text, "Private object pointer"), + f("files", "content_type", Kind::Text, "File metadata"), + f( + "mcp_connectors", + "description", + Kind::Text, + "Connector description", + ), + f( + "mcp_connectors", + "mcp_server_url", + Kind::Text, + "Private endpoint", + ), + f("mcp_connectors", "auth_config", Kind::Json, "Credentials"), + f( + "mcp_connectors", + "error_message", + Kind::Text, + "Upstream error", + ), + f( + "mcp_connectors", + "capabilities", + Kind::Json, + "Private tool schemas", + ), + f( + "mcp_connectors", + "metadata", + Kind::Json, + "Connector metadata", + ), + f( + "mcp_connector_usage", + "request_payload", + Kind::Json, + "Tool request", + ), + f( + "mcp_connector_usage", + "response_payload", + Kind::Json, + "Tool response", + ), + f( + "mcp_connector_usage", + "error_message", + Kind::Text, + "Tool error", + ), +]; + +#[derive(Clone)] +pub struct DatabaseEncryptionState { + pub pool: DbPool, + key: [u8; 32], +} + +impl DatabaseEncryptionState { + pub fn new(pool: DbPool, hex_key: &str) -> anyhow::Result { + let bytes = hex::decode(hex_key)?; + let len = bytes.len(); + let key = bytes + .try_into() + .map_err(|_| anyhow::anyhow!("database encryption key must be 32 bytes, got {len}"))?; + Ok(Self { pool, key }) + } + + pub fn recover_jobs(&self) { + let state = self.clone(); + tokio::spawn(async move { + let result = async { + let client = state.pool.get().await?; + let rows = client + .query( + "SELECT id FROM database_encryption_jobs WHERE status IN ('queued', 'running') ORDER BY created_at", + &[], + ) + .await?; + for row in rows { + spawn_job(state.clone(), row.get(0)); + } + anyhow::Ok(()) + } + .await; + if result.is_err() { + tracing::error!( + error_class = "database_encryption_job_recovery_failed", + "Failed to recover database encryption jobs" + ); + } + }); + } +} + +#[derive(Debug, Deserialize, Serialize, Default)] +pub struct Scope { + #[serde(default)] + tables: Vec, + #[serde(default)] + fields: Vec, +} + +#[derive(Debug, Deserialize, Serialize)] +struct FieldName { + table: String, + column: String, +} + +#[derive(Debug, Deserialize, Default)] +pub struct ScanRequest { + #[serde(default)] + scope: Scope, + limit: Option, + #[serde(default)] + include_approved_plaintext: bool, +} + +#[derive(Debug, Deserialize)] +#[serde(rename_all = "snake_case")] +enum Mode { + DryRun, + Execute, +} + +#[derive(Debug, Deserialize)] +pub struct CreateJobRequest { + mode: Mode, + #[serde(default)] + scope: Scope, + #[serde(default = "batch_default")] + batch_size: i64, + max_rows: Option, + #[serde(default = "actions_default")] + actions: Vec, +} + +fn batch_default() -> i64 { + 100 +} + +fn actions_default() -> Vec { + vec!["encrypt".into()] +} + +#[derive(Debug, Serialize)] +pub struct FieldCount { + table: String, + column: String, + classification: &'static str, + plaintext: i64, + encrypted: i64, + empty: i64, + invalid_envelope: i64, + scanned: i64, + complete: bool, +} + +#[derive(Debug, Serialize)] +pub struct ScanResponse { + run_id: Uuid, + status: &'static str, + fields: Vec, + totals: Value, +} + +#[derive(Debug, Serialize)] +pub struct JobResponse { + job_id: Uuid, + status: String, + mode: String, + scope: Value, + actions: Value, + progress: Value, + cursor: Value, + created_at: DateTime, + completed_at: Option>, + last_error_class: Option, + last_error_message: Option, +} + +type ApiResult = Result, (StatusCode, Json)>; + +// Request validation and database inventory helpers. + +fn bad(m: &str) -> (StatusCode, Json) { + ( + StatusCode::BAD_REQUEST, + Json(json!({"error":{"code":"invalid_request","message":m}})), + ) +} + +fn internal(e: impl std::fmt::Display) -> (StatusCode, Json) { + let _ = e; + tracing::error!( + error_class = "database_encryption", + "database encryption operation failed" + ); + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json( + json!({"error":{"code":"database_encryption_failed","message":"Database encryption operation failed"}}), + ), + ) +} + +fn selected(scope: &Scope) -> Result, (StatusCode, Json)> { + let known_tables = FIELDS + .iter() + .map(|f| f.table) + .collect::>(); + if scope + .tables + .iter() + .any(|t| !known_tables.contains(t.as_str())) + { + return Err(bad("scope contains an unknown table")); + } + if scope.fields.iter().any(|x| { + !FIELDS + .iter() + .any(|f| f.table == x.table && f.column == x.column) + }) { + return Err(bad("scope contains an unknown field")); + } + let v = FIELDS + .iter() + .filter(|f| { + scope.tables.is_empty() && scope.fields.is_empty() + || scope.tables.iter().any(|t| t == f.table) + || scope + .fields + .iter() + .any(|x| x.table == f.table && x.column == f.column) + }) + .collect::>(); + if v.is_empty() { + Err(bad( + "scope does not contain a registered confidential field", + )) + } else { + Ok(v) + } +} + +fn normalize_scope(scope: &mut Scope) { + scope.tables.sort(); + scope.tables.dedup(); + scope + .fields + .sort_by(|left, right| (&left.table, &left.column).cmp(&(&right.table, &right.column))); + scope + .fields + .dedup_by(|left, right| left.table == right.table && left.column == right.column); +} + +fn predicate(f: &Field) -> String { + match f.kind { + Kind::Json => format!( + "{0} IS NOT NULL AND jsonb_typeof({0})='object' \ + AND {0}->>'{MARKER}'='true' AND {0}->>'version'='1' \ + AND {0}->>'alg'='AES-256-GCM' \ + AND {0} ?& ARRAY['key_id','nonce','ciphertext']", + f.column + ), + Kind::Text => format!( + "{0} IS NOT NULL AND {0} LIKE '{{\"{MARKER}\":true,%'", + f.column + ), + } +} + +fn operational_scope(entries: Vec) -> anyhow::Result { + let mut scope = Scope::default(); + for entry in entries { + if let Some((table, column)) = entry.split_once('.') { + scope.fields.push(FieldName { + table: table.to_string(), + column: column.to_string(), + }); + } else { + scope.tables.push(entry); + } + } + normalize_scope(&mut scope); + selected(&scope).map_err(|(_, Json(value))| anyhow::anyhow!("invalid scope: {value}"))?; + Ok(scope) +} + +async fn verify_classification_inventory(state: &DatabaseEncryptionState) -> anyhow::Result { + let client = state.pool.get().await?; + let rows = client + .query( + "SELECT table_name,column_name FROM information_schema.columns \ + WHERE table_schema='public' \ + AND (data_type IN ('text','json','jsonb','character varying','character') OR udt_name='citext') \ + AND table_name <> 'refinery_schema_history' \ + ORDER BY table_name,column_name", + &[], + ) + .await?; + let mut classified = Vec::with_capacity(rows.len()); + let mut missing = Vec::new(); + for row in rows { + let table: String = row.get(0); + let column: String = row.get(1); + match policy_for(&table, &column) { + Some(policy) => classified.push(json!({ + "table": table, + "column": column, + "classification": policy.classification, + "reason": policy.reason, + })), + None => missing.push(format!("{table}.{column}")), + } + } + if !missing.is_empty() { + anyhow::bail!("unclassified text/JSON database columns: {missing:?}"); + } + let encrypted_inventory: std::collections::HashSet<_> = classified + .iter() + .filter(|field| field["classification"] == json!(Classification::Encrypt)) + .filter_map(|field| Some((field["table"].as_str()?, field["column"].as_str()?))) + .collect(); + for field in FIELDS { + if !encrypted_inventory.contains(&(field.table, field.column)) { + anyhow::bail!( + "encryption registry field has no encrypt policy: {}.{}", + field.table, + field.column + ); + } + } + let registry: std::collections::HashSet<_> = FIELDS + .iter() + .map(|field| (field.table, field.column)) + .collect(); + for field in &encrypted_inventory { + if !registry.contains(field) { + anyhow::bail!( + "encrypt policy has no implementation in the encryption registry: {}.{}", + field.0, + field.1 + ); + } + } + Ok(json!({"complete": true, "fields": classified})) +} + +pub async fn operational_scan( + state: &DatabaseEncryptionState, + entries: Vec, +) -> anyhow::Result { + let inventory = verify_classification_inventory(state).await?; + let scope = operational_scope(entries)?; + let fields = selected(&scope).map_err(|_| anyhow::anyhow!("invalid scope"))?; + let counts = counts(state, &fields, None).await?; + Ok(json!({"inventory": inventory, "fields": counts, "totals": totals(&counts)})) +} + +pub async fn operational_verify( + state: &DatabaseEncryptionState, + entries: Vec, +) -> anyhow::Result { + let report = operational_scan(state, entries).await?; + let totals = &report["totals"]; + let pass = totals["plaintext"].as_i64() == Some(0) + && totals["invalid_envelope"].as_i64() == Some(0) + && totals["complete"].as_bool() == Some(true); + Ok(json!({"pass": pass, "report": report})) +} + +pub async fn operational_migrate( + state: &DatabaseEncryptionState, + entries: Vec, + batch_size: i64, + max_rows: Option, + resume: Option, + operator: &str, +) -> anyhow::Result { + if !(1..=1000).contains(&batch_size) { + anyhow::bail!("batch_size must be between 1 and 1000"); + } + if max_rows.is_some_and(|value| value <= 0) { + anyhow::bail!("max_rows must be greater than zero"); + } + let verification_entries = entries.clone(); + let scope = operational_scope(entries)?; + if scope.tables.is_empty() && scope.fields.is_empty() { + anyhow::bail!("an explicit scope is required"); + } + let id = if let Some(id) = resume { + let client = state.pool.get().await?; + let persisted_scope: Value = client + .query_opt( + "SELECT scope FROM database_encryption_jobs \ + WHERE id=$1 AND status IN ('queued','running','failed')", + &[&id], + ) + .await? + .ok_or_else(|| anyhow::anyhow!("job does not exist or is not resumable"))? + .get(0); + if persisted_scope != serde_json::to_value(&scope)? { + anyhow::bail!("resume scope does not match the persisted job scope"); + } + client + .execute( + "UPDATE database_encryption_jobs SET status='queued',completed_at=NULL,last_error_class=NULL,last_error_message=NULL WHERE id=$1", + &[&id], + ) + .await?; + id + } else { + let id = Uuid::new_v4(); + let client = state.pool.get().await?; + client.execute( + "INSERT INTO database_encryption_jobs(id,mode,status,scope,actions,batch_size,max_rows,admin_actor,operator) \ + VALUES($1,'execute','queued',$2,$3,$4,$5,NULL,$6)", + &[&id, &serde_json::to_value(&scope)?, &json!(["encrypt"]), &batch_size, &max_rows, &operator], + ).await?; + id + }; + run_job(state, id).await?; + let client = state.pool.get().await?; + let status: String = client + .query_one( + "SELECT status FROM database_encryption_jobs WHERE id=$1", + &[&id], + ) + .await? + .get(0); + if status != "completed" { + anyhow::bail!("database encryption job ended with status {status}"); + } + let verification = operational_verify(state, verification_entries).await?; + if verification["pass"] != true { + anyhow::bail!("post-migration verification failed"); + } + Ok(id) +} + +async fn counts( + state: &DatabaseEncryptionState, + fields: &[&Field], + limit: Option, +) -> anyhow::Result> { + let mut client = state.pool.get().await?; + let mut out = vec![]; + for f in fields { + let scan_limit = limit.map(|value| value.clamp(1, 100_000)); + let max_id_query = format!("SELECT id FROM {} ORDER BY id DESC LIMIT 1", f.table); + let max_id = client + .query_opt(&max_id_query, &[]) + .await? + .map(|row| row.get::<_, Uuid>(0)); + let mut count = FieldCount { + table: f.table.into(), + column: f.column.into(), + classification: "encrypt", + empty: 0, + encrypted: 0, + plaintext: 0, + invalid_envelope: 0, + scanned: 0, + complete: true, + }; + let Some(max_id) = max_id else { + out.push(count); + continue; + }; + let mut after_id = Uuid::nil(); + + loop { + let remaining = scan_limit + .map(|value| value - count.scanned) + .unwrap_or(1_000); + if remaining <= 0 { + count.complete = false; + break; + } + let page_size = remaining.min(1_000); + let transaction = client.transaction().await?; + transaction + .batch_execute("SET LOCAL statement_timeout = '5s'") + .await?; + let query = format!( + "SELECT id,{}::text FROM {} WHERE id>$1 AND id<=$2 ORDER BY id LIMIT $3", + f.column, f.table + ); + let rows = transaction + .query(&query, &[&after_id, &max_id, &page_size]) + .await?; + transaction.commit().await?; + if rows.is_empty() { + break; + } + for row in &rows { + let id: Uuid = row.get(0); + after_id = id; + let raw: Option = row.get(1); + let Some(raw) = raw else { + count.empty += 1; + continue; + }; + match serde_json::from_str::(&raw) { + Ok(value) if value[MARKER] == true => { + if decrypt_envelope(&state.key, f, id, &raw).is_ok() { + count.encrypted += 1; + } else { + count.invalid_envelope += 1; + } + } + _ => count.plaintext += 1, + } + } + count.scanned += rows.len() as i64; + if after_id == max_id { + break; + } + } + out.push(count); + } + Ok(out) +} + +fn totals(c: &[FieldCount]) -> Value { + json!({"plaintext":c.iter().map(|x|x.plaintext).sum::(),"encrypted":c.iter().map(|x|x.encrypted).sum::(),"empty":c.iter().map(|x|x.empty).sum::(),"invalid_envelope":c.iter().map(|x|x.invalid_envelope).sum::(),"scanned":c.iter().map(|x|x.scanned).sum::(),"complete":c.iter().all(|x|x.complete)}) +} + +// Admin scan and verification endpoints. +pub async fn scan( + Extension(state): Extension, + Extension(_): Extension, + Json(req): Json, +) -> ApiResult { + let fs = selected(&req.scope)?; + let cs = counts(&state, &fs, req.limit).await.map_err(internal)?; + let t = totals(&cs); + let _ = req.include_approved_plaintext; + Ok(Json(ScanResponse { + run_id: Uuid::new_v4(), + status: "completed", + fields: cs, + totals: t, + })) +} + +fn envelope(key: &[u8; 32], f: &Field, id: Uuid, plain: &str) -> anyhow::Result { + database::field_encryption::encrypt(key, f.table, f.column, id, plain) +} + +// Job lifecycle and authenticated envelope helpers. + +fn decrypt_envelope(key: &[u8; 32], f: &Field, id: Uuid, encoded: &str) -> anyhow::Result { + database::field_encryption::decrypt(key, f.table, f.column, id, encoded) +} + +pub async fn create_job( + Extension(state): Extension, + Extension(admin): Extension, + Json(req): Json, +) -> Result<(StatusCode, Json), (StatusCode, Json)> { + if req.scope.tables.is_empty() && req.scope.fields.is_empty() { + return Err(bad( + "an explicit scope is required for database encryption jobs", + )); + } + if !(1..=1000).contains(&req.batch_size) { + return Err(bad("batch_size must be between 1 and 1000")); + } + if req.actions.as_slice() != ["encrypt"] { + return Err(bad("actions must contain exactly one encrypt action")); + } + if req.max_rows.is_some_and(|max_rows| max_rows <= 0) { + return Err(bad("max_rows must be greater than zero")); + } + let mut scope_request = req.scope; + normalize_scope(&mut scope_request); + selected(&scope_request)?; + let id = Uuid::new_v4(); + let mode = match req.mode { + Mode::DryRun => "dry_run", + Mode::Execute => "execute", + }; + let scope = serde_json::to_value(&scope_request).map_err(internal)?; + let actions = json!(req.actions); + let client = state.pool.get().await.map_err(internal)?; + client.execute("INSERT INTO database_encryption_jobs(id,mode,status,scope,actions,batch_size,max_rows,admin_actor) VALUES($1,$2,'queued',$3,$4,$5,$6,$7)",&[&id,&mode,&scope,&actions,&req.batch_size,&req.max_rows,&admin.0.id]).await.map_err(internal)?; + drop(client); + spawn_job(state.clone(), id); + let response = get_inner(&state, id).await?; + Ok((StatusCode::ACCEPTED, response)) +} + +fn spawn_job(state: DatabaseEncryptionState, id: Uuid) { + tokio::spawn(async move { + if run_job(&state, id).await.is_err() { + if let Ok(client) = state.pool.get().await { + let _ = client + .execute( + "UPDATE database_encryption_jobs SET status='failed',last_error_class='batch_failed',last_error_message='batch_failed',completed_at=NOW() WHERE id=$1 AND status IN ('queued','running')", + &[&id], + ) + .await; + } + tracing::error!( + job_id = %id, + error_class = "database_encryption_batch_failed", + "Database encryption job failed" + ); + } + }); +} + +fn advisory_lock_key(id: Uuid) -> i64 { + i64::from_be_bytes( + id.as_bytes()[..8] + .try_into() + .expect("UUID prefix is 8 bytes"), + ) +} + +async fn run_job(state: &DatabaseEncryptionState, id: Uuid) -> anyhow::Result<()> { + let mut client = state.pool.get().await?; + let lock_key = advisory_lock_key(id); + let locked: bool = client + .query_one("SELECT pg_try_advisory_lock($1)", &[&lock_key]) + .await? + .get(0); + if !locked { + return Ok(()); + } + + let result = run_locked_job(state, id, &mut client).await; + let _ = client + .query_one("SELECT pg_advisory_unlock($1)", &[&lock_key]) + .await; + result +} + +async fn run_locked_job( + state: &DatabaseEncryptionState, + id: Uuid, + client: &mut tokio_postgres::Client, +) -> anyhow::Result<()> { + let job = client + .query_opt( + "UPDATE database_encryption_jobs SET status='running',started_at=COALESCE(started_at,NOW()) WHERE id=$1 AND status IN ('queued','running') RETURNING mode,scope,batch_size,max_rows,cursor,progress", + &[&id], + ) + .await?; + let Some(job) = job else { + return Ok(()); + }; + let mode: String = job.get("mode"); + let scope: Scope = serde_json::from_value(job.get("scope"))?; + let fields = selected(&scope).map_err(|_| anyhow::anyhow!("invalid persisted scope"))?; + let batch: i64 = job.get("batch_size"); + let max: Option = job.get("max_rows"); + let cursor: Value = job.get("cursor"); + let progress: Value = job.get("progress"); + let mut field_index = cursor["field_index"].as_u64().unwrap_or(0) as usize; + let mut after_id = cursor["after_id"] + .as_str() + .and_then(|value| Uuid::parse_str(value).ok()) + .unwrap_or(Uuid::nil()); + let mut processed = progress["processed"].as_i64().unwrap_or(0); + let mut encrypted = progress["encrypted"].as_i64().unwrap_or(0); + + while field_index < fields.len() && max.is_none_or(|limit| processed < limit) { + let field = fields[field_index]; + let cap = max + .map(|limit| (limit - processed).min(batch)) + .unwrap_or(batch); + let predicate = predicate(field); + let transaction = client.transaction().await?; + transaction + .batch_execute("SET LOCAL statement_timeout = '30s'") + .await?; + let cancelled: bool = transaction + .query_one( + "SELECT cancel_requested_at IS NOT NULL FROM database_encryption_jobs WHERE id=$1", + &[&id], + ) + .await? + .get(0); + if cancelled { + transaction + .execute( + "UPDATE database_encryption_jobs SET status='cancelled',completed_at=NOW() WHERE id=$1", + &[&id], + ) + .await?; + transaction.commit().await?; + return Ok(()); + } + + let locking_clause = if mode == "execute" { " FOR UPDATE" } else { "" }; + let query = format!( + "SELECT id,{0}::text FROM {1} WHERE id>$1 AND {0} IS NOT NULL AND NOT({predicate}) ORDER BY id LIMIT $2{locking_clause}", + field.column, field.table + ); + let rows = transaction.query(&query, &[&after_id, &cap]).await?; + if rows.is_empty() { + field_index += 1; + after_id = Uuid::nil(); + } else { + let mut row_ids = Vec::with_capacity(rows.len()); + let mut encrypted_values = Vec::with_capacity(rows.len()); + for row in &rows { + let row_id: Uuid = row.get(0); + after_id = row_id; + if mode == "execute" { + let plaintext: String = row.get(1); + row_ids.push(row_id); + encrypted_values.push(envelope(&state.key, field, row_id, &plaintext)?); + } + } + if mode == "execute" { + let value_expression = match field.kind { + Kind::Json => "batch.value::jsonb", + Kind::Text => "batch.value", + }; + let update = format!( + "UPDATE {table} AS target SET {column}={value_expression} \ + FROM UNNEST($1::uuid[], $2::text[]) AS batch(id,value) \ + WHERE target.id=batch.id", + table = field.table, + column = field.column, + ); + encrypted += transaction + .execute(&update, &[&row_ids, &encrypted_values]) + .await? as i64; + } + processed += rows.len() as i64; + } + + transaction + .execute( + "UPDATE database_encryption_jobs SET progress=$2,cursor=$3 WHERE id=$1", + &[ + &id, + &json!({"processed":processed,"encrypted":encrypted}), + &json!({"field_index":field_index,"after_id":after_id}), + ], + ) + .await?; + transaction.commit().await?; + } + + client + .execute( + "UPDATE database_encryption_jobs SET status='completed',completed_at=NOW() WHERE id=$1", + &[&id], + ) + .await?; + Ok(()) +} + +pub async fn get_job( + Extension(state): Extension, + Extension(_): Extension, + Path(id): Path, +) -> ApiResult { + get_inner(&state, id).await +} + +async fn get_inner(state: &DatabaseEncryptionState, id: Uuid) -> ApiResult { + let c = state.pool.get().await.map_err(internal)?; + let r=c.query_opt("SELECT id,status,mode,scope,actions,progress,cursor,created_at,completed_at,last_error_class,last_error_message FROM database_encryption_jobs WHERE id=$1",&[&id]).await.map_err(internal)?.ok_or_else(||(StatusCode::NOT_FOUND,Json(json!({"error":{"code":"job_not_found","message":"Database encryption job not found"}}))))?; + Ok(Json(JobResponse { + job_id: r.get(0), + status: r.get(1), + mode: r.get(2), + scope: r.get(3), + actions: r.get(4), + progress: r.get(5), + cursor: r.get(6), + created_at: r.get(7), + completed_at: r.get(8), + last_error_class: r.get(9), + last_error_message: r.get(10), + })) +} + +pub async fn cancel_job( + Extension(state): Extension, + Extension(_): Extension, + Path(id): Path, +) -> ApiResult { + let c = state.pool.get().await.map_err(internal)?; + let n=c.execute("UPDATE database_encryption_jobs SET cancel_requested_at=NOW(),status=CASE WHEN status='queued' THEN 'cancelled' ELSE status END,completed_at=CASE WHEN status='queued' THEN NOW() ELSE completed_at END WHERE id=$1 AND status IN('queued','running')",&[&id]).await.map_err(internal)?; + if n == 0 { + return Err(bad("job is not cancellable")); + } + get_inner(&state, id).await +} + +#[derive(Debug, Deserialize, Default)] +pub struct VerifyRequest { + #[serde(default)] + scope: Scope, + #[serde(default = "yes")] + fail_on_approved_plaintext_without_reason: bool, +} + +fn yes() -> bool { + true +} + +#[derive(Debug, Serialize)] +pub struct VerifyResponse { + pass: bool, + fields: Vec, + failing_fields: Vec, +} + +pub async fn verify( + Extension(state): Extension, + Extension(_): Extension, + Json(req): Json, +) -> ApiResult { + let fs = selected(&req.scope)?; + let cs = counts(&state, &fs, None).await.map_err(internal)?; + let fails = cs + .iter() + .filter(|x| x.plaintext > 0 || x.invalid_envelope > 0) + .map(|x| { + let reason = if x.invalid_envelope > 0 { + "invalid_envelope" + } else { + "plaintext_remaining" + }; + json!({"table":x.table,"column":x.column,"reason_code":reason}) + }) + .collect::>(); + let _ = req.fail_on_approved_plaintext_without_reason; + Ok(Json(VerifyResponse { + pass: fails.is_empty(), + fields: cs, + failing_fields: fails, + })) +} + +#[cfg(test)] +mod tests { + use super::*; + #[test] + fn envelope_hides_plaintext() { + let v = envelope(&[7; 32], &FIELDS[0], Uuid::nil(), "secret").unwrap(); + assert!(v.contains(MARKER)); + assert!(!v.contains("secret")); + assert_eq!( + decrypt_envelope(&[7; 32], &FIELDS[0], Uuid::nil(), &v).unwrap(), + "secret" + ); + } + #[test] + fn envelope_authenticates_context_and_key() { + let v = envelope(&[7; 32], &FIELDS[0], Uuid::nil(), "secret").unwrap(); + assert!(decrypt_envelope(&[8; 32], &FIELDS[0], Uuid::nil(), &v).is_err()); + assert!(decrypt_envelope(&[7; 32], &FIELDS[1], Uuid::nil(), &v).is_err()); + assert!(decrypt_envelope(&[7; 32], &FIELDS[0], Uuid::new_v4(), &v).is_err()); + } + + #[test] + fn malformed_envelope_is_rejected() { + let malformed = json!({MARKER: true, "version": 1, "alg": "AES-256-GCM"}).to_string(); + assert!(decrypt_envelope(&[7; 32], &FIELDS[0], Uuid::nil(), &malformed).is_err()); + } + #[test] + fn key_requires_32_decoded_bytes() { + assert!(hex::decode("not-hex").is_err()); + let short: Result<[u8; 32], _> = hex::decode("00").unwrap().try_into(); + assert!(short.is_err()); + } + #[test] + fn registry_unique() { + let mut n = FIELDS + .iter() + .map(|f| (f.table, f.column)) + .collect::>(); + n.sort(); + n.dedup(); + assert_eq!(n.len(), FIELDS.len()); + } +} diff --git a/crates/api/src/database_encryption_inventory.rs b/crates/api/src/database_encryption_inventory.rs new file mode 100644 index 000000000..d8ba35df5 --- /dev/null +++ b/crates/api/src/database_encryption_inventory.rs @@ -0,0 +1,165 @@ +use serde::Serialize; + +#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize)] +#[serde(rename_all = "snake_case")] +pub enum Classification { + Encrypt, + Remove, + ApprovedPlaintext, +} + +#[derive(Clone, Copy, Debug, Serialize)] +pub struct Policy { + pub classification: Classification, + pub reason: &'static str, +} + +const fn encrypt(reason: &'static str) -> Policy { + Policy { + classification: Classification::Encrypt, + reason, + } +} + +const fn approved(reason: &'static str) -> Policy { + Policy { + classification: Classification::ApprovedPlaintext, + reason, + } +} + +const HASH: &str = "One-way digest or non-secret prefix required for indexed lookup"; +const ENUM: &str = "Bounded operational enum required for filtering and constraints"; +const ID: &str = "Operational identifier required for indexed lookup, routing, or audit"; +const DISPLAY: &str = "Queryable display metadata approved pending searchable-encryption support"; +const CONFIG: &str = "Queryable operational configuration required by the service"; +const AUDIT: &str = "Restricted operator audit metadata; must not contain customer payloads"; +const ERROR: &str = "Redacted operational error; producers must not include customer payloads"; + +/// Explicit policy for every application-owned text, varchar, char, JSON, and JSONB column. +/// Unknown columns fail the worker's inventory gate. +pub fn policy_for(table: &str, column: &str) -> Option { + let policy = match (table, column) { + ("response_items", "item") => encrypt("Persisted customer transcript and tool payload"), + ("responses", "instructions") => encrypt("Customer prompt instructions"), + ("responses", "metadata") => encrypt("Customer response metadata"), + ("conversations", "metadata") => encrypt("Customer conversation title and metadata"), + ("files", "filename" | "content_type") => encrypt("Customer file metadata"), + ("files", "storage_key") => encrypt("Private object pointer"), + ("mcp_connectors", "description" | "mcp_server_url" | "auth_config" | "error_message" | "capabilities" | "metadata") => encrypt("Customer MCP configuration or payload"), + ("mcp_connector_usage", "request_payload" | "response_payload" | "error_message") => encrypt("Customer MCP request, response, or error"), + + ("admin_access_token", "token_hash") | ("api_keys", "key_hash" | "key_prefix") + | ("organization_reporting_tokens", "token_hash" | "token_prefix") + | ("refresh_tokens", "token_hash") => approved(HASH), + ("admin_access_token", "name" | "creation_reason" | "revocation_reason" | "user_agent") + | ("aml_allowlisted_accounts", "reason") + | ("organization_limits_history", "changed_by" | "change_reason" | "changed_by_user_email") + | ("model_history", "change_reason" | "changed_by_user_email") + | ("model_deprecation_email_deliveries", "initiated_by_user_email") + | ("model_pricing_change_email_deliveries", "initiated_by_user_email") + | ("scheduled_model_pricing_changes", "cancelled_by_user_email" | "created_by_user_email" | "change_reason") => approved(AUDIT), + ("aml_allowlisted_accounts", "account_id") + | ("aml_reports", "account_id" | "report_id") + | ("chat_signatures", "chat_id") + | ("model_aliases", "alias_name") + | ("model_deprecation_email_deliveries", "model_name" | "successor_model_name" | "email_message_id") + | ("model_pricing_change_email_deliveries", "email_message_id") + | ("model_history", "model_name" | "hugging_face_id" | "openrouter_slug") + | ("models", "model_name" | "hugging_face_id" | "openrouter_slug") + | ("organization_staking_farm_sources", "near_account_id" | "network_id" | "contract_id" | "farm_product_id" | "farm_price_id") + | ("organization_usage_log", "model_name" | "provider_request_id") + | ("scheduled_model_pricing_changes", "model_name") + | ("services", "service_name") => approved(ID), + ("aml_allowlisted_accounts", "address_type") + | ("aml_reports", "flow" | "provider" | "address_type" | "risk_level") + | ("chat_signatures", "signing_algo" | "signature_kind") + | ("feature_request_targets", "kind" | "status") + | ("feature_request_votes", "source") + | ("files", "purpose") + | ("mcp_connector_usage", "method") + | ("mcp_connectors", "auth_type" | "connection_status") + | ("model_deprecation_email_deliveries", "status") + | ("model_pricing_change_email_deliveries", "status") + | ("model_history", "provider_type" | "quantization" | "attestation_policy") + | ("models", "provider_type" | "quantization" | "attestation_policy") + | ("oauth_states", "provider") + | ("organization_invitations", "role" | "status" | "email_status") + | ("organization_limits_history", "credit_type" | "source" | "currency") + | ("organization_members", "role") + | ("organization_staking_farm_sources", "status" | "sync_status") + | ("organization_usage_log", "request_type" | "inference_type" | "stop_reason" | "served_provider_tier" | "served_provider_type" | "service_tier" | "context_band") + | ("responses", "status") + | ("scheduled_model_pricing_changes", "status") + | ("services", "unit") + | ("users", "auth_provider") => approved(ENUM), + ("feature_request_targets", "key") => approved(ID), + ("feature_request_targets", "title") + | ("feature_request_votes", "note") + | ("mcp_connectors", "name") + | ("model_deprecation_email_deliveries", "model_display_name" | "organization_name") + | ("model_pricing_change_email_deliveries", "organization_name") + | ("model_history", "model_display_name" | "model_description" | "owned_by") + | ("models", "model_display_name" | "model_description" | "owned_by") + | ("organizations", "name" | "description") + | ("scheduled_model_pricing_changes", "model_display_name") + | ("services", "display_name" | "description") + | ("users", "display_name" | "avatar_url") + | ("workspaces", "name" | "description") + | ("api_keys", "name") + | ("organization_reporting_tokens", "name") => approved(DISPLAY), + ("model_history", "model_icon" | "provider_config" | "input_modalities" | "output_modalities" | "inference_url" | "text_pricing") + | ("models", "model_icon" | "provider_config" | "input_modalities" | "output_modalities" | "inference_url" | "text_pricing") + | ("organizations", "settings") + | ("scheduled_model_pricing_changes", "old_text_pricing" | "new_text_pricing") + | ("workspaces", "settings") => approved(CONFIG), + ("aml_reports", "reason") + | ("model_deprecation_email_deliveries", "email_last_error") + | ("model_pricing_change_email_deliveries", "email_last_error") + | ("organization_invitations", "email_last_error") + | ("organization_staking_farm_sources", "last_sync_error") + | ("scheduled_model_pricing_changes", "last_error") => approved(ERROR), + ("aml_reports", "result_json") => approved("Restricted compliance result required for enforcement and audit"), + ("chat_signatures", "text") => approved("Cryptographic attestation statement returned to its tenant"), + ("chat_signatures", "signature" | "signing_address") => approved("Public cryptographic verification material"), + ("database_encryption_jobs", "mode" | "status" | "scope" | "actions" | "cursor" | "progress" | "last_error_class" | "last_error_message" | "operator") => approved("Encryption-worker control state containing identifiers, counters, and redacted errors only"), + ("model_deprecation_email_deliveries", "recipient_email") + | ("model_pricing_change_email_deliveries", "recipient_email") => approved("Delivery address required for notification audit"), + ("near_used_nonces", "nonce_hex") => approved("Public replay-prevention nonce"), + ("oauth_states", "state") => approved("Short-lived random correlation token required for indexed OAuth lookup"), + ("oauth_states", "pkce_verifier" | "frontend_callback") => approved("Short-lived OAuth protocol value; expired state rows are deleted"), + ("organization_invitations", "email") => approved("Invitation identity required for indexed lookup and delivery"), + ("organization_invitations", "token") => approved("Short-lived bearer token required for indexed invitation acceptance; hash migration is required separately"), + ("organization_invitations", "email_message_id") => approved(ID), + ("organization_staking_farm_sources", "active_positions") => approved("Queryable staking accounting state"), + ("organization_usage_log", "billing_details") => approved("Queryable billing ledger inputs without prompt or response content"), + ("refresh_tokens", "ip_address" | "user_agent") => approved("Restricted account-security audit attribute"), + ("responses", "model") => approved(ID), + ("responses", "usage") => approved("Numeric token accounting without customer content"), + ("responses", "next_response_ids") => approved("Response graph identifiers"), + ("users", "email") => approved("Queryable login identity protected by account access controls"), + ("users", "username") => approved("Queryable unique account identity"), + ("users", "provider_user_id") => approved(ID), + _ => return None, + }; + Some(policy) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn sensitive_examples_have_explicit_policies() { + for (table, column) in [ + ("oauth_states", "pkce_verifier"), + ("organization_invitations", "token"), + ("users", "email"), + ("organizations", "settings"), + ("workspaces", "settings"), + ] { + let policy = policy_for(table, column).expect("field must be classified"); + assert!(!policy.reason.is_empty()); + } + } +} diff --git a/crates/api/src/lib.rs b/crates/api/src/lib.rs index 374dc8673..4d3b7985b 100644 --- a/crates/api/src/lib.rs +++ b/crates/api/src/lib.rs @@ -1,5 +1,7 @@ pub mod consts; pub mod conversions; +pub mod database_encryption; +pub mod database_encryption_inventory; pub mod middleware; pub mod models; pub mod ohttp_gateway; @@ -319,6 +321,9 @@ pub async fn init_domain_services_with_pool( inference_provider_pool: Arc, metrics_service: Arc, ) -> DomainServices { + if let Ok(key) = database::field_encryption::parse_key(&config.s3.encryption_key) { + database.pool().set_encryption_key(key); + } // Give the provider pool the metrics sink so it can emit the per-tier / // fallback counter (cloud_api.provider.requests) from the one layer that // knows which trust tier served each request. @@ -2189,14 +2194,14 @@ pub fn build_admin_routes( usage_service: services.usage_service, staking_farm_service: services.staking_farm_service, aml_service: services.aml_service, - config, + config: config.clone(), admin_access_token_repository, inference_provider_pool: services.inference_provider_pool, github_dispatcher, infra_service, }; - Router::new() + let admin_routes = Router::new() .route( "/admin/models", axum::routing::get(admin_list_models).patch(batch_upsert_models), @@ -2354,7 +2359,9 @@ pub fn build_admin_routes( "/admin/access-tokens/{token_id}", axum::routing::delete(delete_admin_access_token), ) - .with_state(admin_app_state) + .with_state(admin_app_state); + + admin_routes // Admin middleware handles both authentication and authorization .layer(from_fn_with_state( auth_state_middleware.clone(), diff --git a/crates/api/tests/e2e_all/client_disconnect.rs b/crates/api/tests/e2e_all/client_disconnect.rs index 0509050b1..59ab79425 100644 --- a/crates/api/tests/e2e_all/client_disconnect.rs +++ b/crates/api/tests/e2e_all/client_disconnect.rs @@ -20,14 +20,25 @@ async fn get_assistant_item_from_db( let rows = client .query( - "SELECT item FROM response_items WHERE conversation_id = $1 ORDER BY created_at DESC", + "SELECT id, item FROM response_items WHERE conversation_id = $1 ORDER BY created_at DESC", &[&conv_uuid], ) .await .expect("Failed to query response_items"); for row in rows { - let item: serde_json::Value = row.get("item"); + let id: uuid::Uuid = row.get("id"); + let mut item: serde_json::Value = row.get("item"); + if let Some(key) = pool.encryption_key() { + item = database::field_encryption::decrypt_json_if_encrypted( + &key, + "response_items", + "item", + id, + item, + ) + .expect("Failed to decrypt response item"); + } if item.get("role").and_then(|v| v.as_str()) == Some("assistant") { return Some(item); } @@ -277,8 +288,13 @@ async fn test_signature_returns_stream_disconnected_on_client_disconnect() { ) .await .expect("Failed to delete signature"); + drop(client); // Update the stop_reason to client_disconnect + // Reacquire after the deliberately disconnected stream: the connection + // used by the request can be closed asynchronously while this test is + // preparing the assertion. + let client = pool.get().await.expect("Failed to get database connection"); client .execute( "UPDATE organization_usage_log SET stop_reason = 'client_disconnect' WHERE response_id = $1", diff --git a/crates/api/tests/e2e_all/conversations.rs b/crates/api/tests/e2e_all/conversations.rs index e079e42f0..531f4713f 100644 --- a/crates/api/tests/e2e_all/conversations.rs +++ b/crates/api/tests/e2e_all/conversations.rs @@ -273,7 +273,7 @@ async fn test_response_stream_fails_with_failed_event_when_inference_fails_at_st let item_rows = client .query( - "SELECT item FROM response_items WHERE conversation_id = $1 ORDER BY created_at ASC", + "SELECT id, item FROM response_items WHERE conversation_id = $1 ORDER BY created_at ASC", &[&conv_uuid], ) .await @@ -281,7 +281,18 @@ async fn test_response_stream_fails_with_failed_event_when_inference_fails_at_st let assistant_items: Vec = item_rows .into_iter() .filter_map(|row| { - let item: serde_json::Value = row.get("item"); + let id: uuid::Uuid = row.get("id"); + let mut item: serde_json::Value = row.get("item"); + if let Some(key) = pool.encryption_key() { + item = database::field_encryption::decrypt_json_if_encrypted( + &key, + "response_items", + "item", + id, + item, + ) + .expect("decrypt response item"); + } if item.get("role").and_then(|v| v.as_str()) == Some("assistant") { Some(item) } else { @@ -2463,7 +2474,7 @@ async fn test_clone_conversation() { #[tokio::test] async fn test_clone_conversation_with_responses_and_items() { - let server = setup_test_server().await; + let (server, database) = setup_test_server_with_database().await; let model_id = setup_qwen_model(&server).await; let org = setup_org_with_credits(&server, 10000000000i64).await; // $10.00 USD let api_key = get_api_key_for_org(&server, org.id).await; @@ -2505,6 +2516,21 @@ async fn test_clone_conversation_with_responses_and_items() { .await; println!("Created response 2: {}", response2.id); + let confidential_response = server + .post("/v1/responses") + .add_header("Authorization", format!("Bearer {api_key}")) + .json(&serde_json::json!({ + "conversation": {"id": original_conv.id}, + "input": "Response with confidential fields", + "instructions": "private system instruction", + "metadata": {"private": "response metadata"}, + "max_output_tokens": 50, + "stream": false, + "model": model_id + })) + .await; + assert_eq!(confidential_response.status_code(), 200); + // Get original conversation items count let original_items = list_conversation_items(&server, original_conv.id.clone(), api_key.clone()).await; @@ -2533,6 +2559,50 @@ async fn test_clone_conversation_with_responses_and_items() { Some("Original Conversation with Messages (Copy)") ); + let original_uuid = uuid::Uuid::parse_str( + original_conv + .id + .strip_prefix("conv_") + .unwrap_or(&original_conv.id), + ) + .expect("original conversation UUID"); + let cloned_uuid = uuid::Uuid::parse_str( + cloned_conv + .id + .strip_prefix("conv_") + .unwrap_or(&cloned_conv.id), + ) + .expect("cloned conversation UUID"); + let client = database.pool().get().await.expect("database connection"); + for conversation_id in [original_uuid, cloned_uuid] { + let metadata: serde_json::Value = client + .query_one( + "SELECT metadata FROM conversations WHERE id=$1", + &[&conversation_id], + ) + .await + .expect("stored conversation metadata") + .get(0); + assert_eq!(metadata[database::field_encryption::MARKER], true); + + let response_row = client + .query_one( + "SELECT instructions,metadata FROM responses WHERE conversation_id=$1 AND instructions IS NOT NULL LIMIT 1", + &[&conversation_id], + ) + .await + .expect("stored confidential response fields"); + assert!(response_row + .get::<_, String>("instructions") + .contains(database::field_encryption::MARKER)); + assert_eq!( + response_row.get::<_, serde_json::Value>("metadata") + [database::field_encryption::MARKER], + true + ); + } + drop(client); + // Get cloned conversation items let cloned_items = list_conversation_items(&server, cloned_conv.id.clone(), api_key.clone()).await; diff --git a/crates/api/tests/e2e_all/database_encryption.rs b/crates/api/tests/e2e_all/database_encryption.rs new file mode 100644 index 000000000..e722377cb --- /dev/null +++ b/crates/api/tests/e2e_all/database_encryption.rs @@ -0,0 +1,108 @@ +use crate::common::db_setup::create_test_pool; +use api::database_encryption::{ + operational_migrate, operational_scan, operational_verify, DatabaseEncryptionState, +}; +use uuid::Uuid; + +#[tokio::test] +async fn operational_worker_encrypts_and_verifies_a_scoped_field() { + let pool = create_test_pool().await; + let key = [42_u8; 32]; + pool.set_encryption_key(key); + let client = pool.get().await.expect("database connection"); + let organization_id = Uuid::new_v4(); + let user_id = Uuid::new_v4(); + let workspace_id = Uuid::new_v4(); + let api_key_id = Uuid::new_v4(); + let file_id = Uuid::new_v4(); + client + .execute( + "INSERT INTO organizations(id,name) VALUES($1,$2)", + &[&organization_id, &format!("worker-org-{organization_id}")], + ) + .await + .expect("organization"); + client + .execute( + "INSERT INTO users(id,email,username,auth_provider,provider_user_id) VALUES($1,$2,$3,'mock',$4)", + &[&user_id, &format!("{user_id}@test.invalid"), &format!("worker-{user_id}"), &user_id.to_string()], + ) + .await + .expect("user"); + client + .execute( + "INSERT INTO workspaces(id,name,organization_id,created_by_user_id) VALUES($1,$2,$3,$4)", + &[&workspace_id, &format!("worker-{workspace_id}"), &organization_id, &user_id], + ) + .await + .expect("workspace"); + client + .execute( + "INSERT INTO api_keys(id,key_hash,name,workspace_id,created_by_user_id,key_prefix) VALUES($1,$2,'worker',$3,$4,'sk-test')", + &[ + &api_key_id, + &format!("{:0>64}", api_key_id.simple()), + &workspace_id, + &user_id, + ], + ) + .await + .expect("API key"); + client + .execute( + "INSERT INTO files(id,filename,bytes,content_type,purpose,storage_key,workspace_id,uploaded_by_api_key_id) \ + VALUES($1,'legacy secret.txt',1,'text/plain','assistants','legacy/key',$2,$3)", + &[&file_id, &workspace_id, &api_key_id], + ) + .await + .expect("legacy plaintext file"); + drop(client); + + let state = + DatabaseEncryptionState::new(pool.clone(), &hex::encode(key)).expect("worker state"); + let job_id = operational_migrate( + &state, + vec!["files.filename".to_string()], + 1, + None, + None, + "e2e-test", + ) + .await + .expect("worker migration"); + + let client = pool.get().await.expect("database connection"); + let row = client + .query_one( + "SELECT f.filename,j.status,j.operator FROM files f CROSS JOIN database_encryption_jobs j \ + WHERE f.id=$1 AND j.id=$2", + &[&file_id, &job_id], + ) + .await + .expect("stored migration result"); + let stored: String = row.get(0); + let status: String = row.get(1); + let operator: String = row.get(2); + assert!(stored.contains(database::field_encryption::MARKER)); + assert!(!stored.contains("legacy secret.txt")); + assert_eq!(status, "completed"); + assert_eq!(operator, "e2e-test"); + + let verification = operational_verify(&state, vec!["files.filename".to_string()]) + .await + .expect("worker verification"); + assert_eq!(verification["pass"], true); +} + +#[tokio::test] +async fn worker_inventory_classifies_the_live_schema() { + let pool = create_test_pool().await; + let key = [43_u8; 32]; + pool.set_encryption_key(key); + let state = DatabaseEncryptionState::new(pool, &hex::encode(key)).expect("worker state"); + let scan = operational_scan(&state, vec!["responses.metadata".to_string()]) + .await + .expect("classification scan"); + assert_eq!(scan["inventory"]["complete"], true); + assert!(scan["inventory"]["fields"].as_array().unwrap().len() > 100); +} diff --git a/crates/api/tests/e2e_all/main.rs b/crates/api/tests/e2e_all/main.rs index 790ac6958..808b3c4d8 100644 --- a/crates/api/tests/e2e_all/main.rs +++ b/crates/api/tests/e2e_all/main.rs @@ -34,6 +34,7 @@ mod concurrent_limit; mod conversations; mod credit_types; mod cross_workspace; +mod database_encryption; mod deser_error_envelope; mod duplicate_names; mod embeddings; diff --git a/crates/api/tests/e2e_all/repositories.rs b/crates/api/tests/e2e_all/repositories.rs index de5467ce2..0a932298b 100644 --- a/crates/api/tests/e2e_all/repositories.rs +++ b/crates/api/tests/e2e_all/repositories.rs @@ -206,6 +206,17 @@ mod response_item_workspace_scoping { .expect("owner listing should succeed"); assert_eq!(own_items.len(), 3, "owner should see all 3 items"); + let client = database.pool().get().await.expect("database connection"); + let stored: serde_json::Value = client + .query_one( + "SELECT item FROM response_items WHERE conversation_id=$1 LIMIT 1", + &[&ws_a.conversation_id.0], + ) + .await + .expect("stored response item") + .get(0); + assert_eq!(stored[database::field_encryption::MARKER], true); + // The same conversation queried with a foreign workspace returns // nothing, even though the conversation ID is known. let foreign_items = repo @@ -312,3 +323,220 @@ mod response_item_workspace_scoping { assert!(foreign.is_none(), "foreign workspace must not see the item"); } } + +// ============================================ +// Confidential repository fields at rest +// ============================================ + +mod database_encryption_at_rest { + use crate::common::*; + use database::models::{CreateMcpConnectorRequest, McpAuthType}; + use database::repositories::{FileRepository, McpConnectorRepository}; + use services::files::ports::CreateFileParams; + use uuid::Uuid; + + #[tokio::test] + async fn file_repository_encrypts_storage_and_returns_plaintext() { + let (server, database) = setup_test_server_with_database().await; + let organization = create_org(&server).await; + let _api_key = get_api_key_for_org(&server, organization.id.clone()).await; + let organization_id = Uuid::parse_str(&organization.id).expect("organization UUID"); + let client = database.pool().get().await.expect("database connection"); + let workspace_id: Uuid = client + .query_one( + "SELECT id FROM workspaces WHERE organization_id=$1 LIMIT 1", + &[&organization_id], + ) + .await + .expect("workspace") + .get(0); + let api_key_id: Uuid = client + .query_one( + "SELECT id FROM api_keys WHERE workspace_id=$1 LIMIT 1", + &[&workspace_id], + ) + .await + .expect("API key") + .get(0); + drop(client); + + let repository = FileRepository::new(database.pool().clone()); + let created = repository + .create(CreateFileParams { + filename: "private.txt".into(), + bytes: 7, + content_type: "text/plain".into(), + purpose: "assistants".into(), + storage_key: "private/object/key".into(), + workspace_id, + uploaded_by_api_key_id: api_key_id, + expires_at: None, + }) + .await + .expect("create encrypted file row"); + assert_eq!(created.filename, "private.txt"); + assert_eq!(created.storage_key, "private/object/key"); + + let client = database.pool().get().await.expect("database connection"); + let row = client + .query_one( + "SELECT filename,content_type,storage_key FROM files WHERE id=$1", + &[&created.id], + ) + .await + .expect("stored file row"); + for column in ["filename", "content_type", "storage_key"] { + let stored: String = row.get(column); + assert!(stored.contains(database::field_encryption::MARKER)); + assert!(!stored.contains("private/object/key")); + } + + let fetched = repository + .get_by_id(created.id) + .await + .expect("read file") + .expect("file exists"); + assert_eq!(fetched.filename, "private.txt"); + assert_eq!(fetched.content_type, "text/plain"); + assert_eq!(fetched.storage_key, "private/object/key"); + + let client = database.pool().get().await.expect("database connection"); + client + .execute( + "UPDATE files SET filename='legacy.txt',content_type='text/legacy',storage_key='legacy/key' WHERE id=$1", + &[&created.id], + ) + .await + .expect("write legacy plaintext row"); + drop(client); + let legacy = repository + .get_by_id(created.id) + .await + .expect("read legacy file") + .expect("legacy file exists"); + assert_eq!(legacy.filename, "legacy.txt"); + assert_eq!(legacy.content_type, "text/legacy"); + assert_eq!(legacy.storage_key, "legacy/key"); + } + + #[tokio::test] + async fn mcp_repository_encrypts_configuration_and_usage_payloads() { + let (server, database) = setup_test_server_with_database().await; + let organization = create_org(&server).await; + let organization_id = Uuid::parse_str(&organization.id).expect("organization UUID"); + let user_id = Uuid::parse_str(MOCK_USER_ID).expect("user UUID"); + let repository = McpConnectorRepository::new(database.pool().clone()); + + let connector = repository + .create( + organization_id, + user_id, + CreateMcpConnectorRequest { + name: format!("private-{}", Uuid::new_v4()), + description: Some("internal tools".into()), + mcp_server_url: "https://internal.example/mcp".into(), + auth_type: McpAuthType::Bearer, + bearer_token: Some("bearer-secret".into()), + }, + ) + .await + .expect("create MCP connector"); + assert_eq!(connector.mcp_server_url, "https://internal.example/mcp"); + assert_eq!( + connector.auth_config.as_ref().unwrap()["token"], + "bearer-secret" + ); + + repository + .log_usage( + connector.id, + user_id, + "tools/call".into(), + Some(serde_json::json!({"argument":"private request"})), + Some(serde_json::json!({"result":"private response"})), + Some(200), + Some("private error".into()), + Some(5), + ) + .await + .expect("log MCP usage"); + + let client = database.pool().get().await.expect("database connection"); + let connector_row = client + .query_one( + "SELECT description,mcp_server_url,auth_config FROM mcp_connectors WHERE id=$1", + &[&connector.id], + ) + .await + .expect("stored connector"); + assert!(connector_row + .get::<_, String>("mcp_server_url") + .contains(database::field_encryption::MARKER)); + assert_eq!( + connector_row.get::<_, serde_json::Value>("auth_config") + [database::field_encryption::MARKER], + true + ); + + let usage_row = client + .query_one( + "SELECT request_payload,response_payload,error_message FROM mcp_connector_usage WHERE connector_id=$1 ORDER BY created_at DESC LIMIT 1", + &[&connector.id], + ) + .await + .expect("stored usage"); + assert_eq!( + usage_row.get::<_, serde_json::Value>("request_payload") + [database::field_encryption::MARKER], + true + ); + assert_eq!( + usage_row.get::<_, serde_json::Value>("response_payload") + [database::field_encryption::MARKER], + true + ); + assert!(usage_row + .get::<_, String>("error_message") + .contains(database::field_encryption::MARKER)); + drop(client); + + let fetched = repository + .get_by_id(connector.id) + .await + .expect("read connector") + .expect("connector exists"); + assert_eq!(fetched.mcp_server_url, "https://internal.example/mcp"); + assert_eq!(fetched.auth_config.unwrap()["token"], "bearer-secret"); + let usage = repository + .get_usage_logs(connector.id, 1) + .await + .expect("read usage"); + assert_eq!( + usage[0].request_payload.as_ref().unwrap()["argument"], + "private request" + ); + assert_eq!( + usage[0].response_payload.as_ref().unwrap()["result"], + "private response" + ); + assert_eq!(usage[0].error_message.as_deref(), Some("private error")); + + let client = database.pool().get().await.expect("database connection"); + client + .execute( + "UPDATE mcp_connectors SET description='legacy description',mcp_server_url='https://legacy.example/mcp',auth_config=$2 WHERE id=$1", + &[&connector.id, &serde_json::json!({"token":"legacy token"})], + ) + .await + .expect("write legacy connector values"); + drop(client); + let legacy = repository + .get_by_id(connector.id) + .await + .expect("read legacy connector") + .expect("legacy connector exists"); + assert_eq!(legacy.description.as_deref(), Some("legacy description")); + assert_eq!(legacy.mcp_server_url, "https://legacy.example/mcp"); + assert_eq!(legacy.auth_config.unwrap()["token"], "legacy token"); + } +} diff --git a/crates/database/Cargo.toml b/crates/database/Cargo.toml index 7d5f92c15..21f68c94c 100644 --- a/crates/database/Cargo.toml +++ b/crates/database/Cargo.toml @@ -51,6 +51,8 @@ refinery = { version = "0.9", features = ["tokio-postgres"] } # For hashing API keys and session tokens sha2 = "0.11" hex = "0.4" +aes-gcm = "0.10" +base64 = "0.22" rand = "0.10.1" # For User-Agent normalization diff --git a/crates/database/src/field_encryption.rs b/crates/database/src/field_encryption.rs new file mode 100644 index 000000000..3228df297 --- /dev/null +++ b/crates/database/src/field_encryption.rs @@ -0,0 +1,204 @@ +use aes_gcm::{ + aead::{rand_core::RngCore, Aead, KeyInit, OsRng, Payload}, + Aes256Gcm, Nonce, +}; +use anyhow::{anyhow, ensure, Context, Result}; +use base64::{engine::general_purpose::STANDARD as BASE64, Engine}; +use serde_json::{json, Value}; +use uuid::Uuid; + +pub const MARKER: &str = "__near_db_encrypted"; + +pub fn is_envelope(value: &Value) -> bool { + value[MARKER] == true + && value["version"] == 1 + && value["alg"] == "AES-256-GCM" + && value["key_id"].as_str().is_some() + && value["nonce"].as_str().is_some() + && value["ciphertext"].as_str().is_some() +} + +pub fn parse_key(hex_key: &str) -> Result<[u8; 32]> { + let bytes = hex::decode(hex_key).context("database encryption key must be hex encoded")?; + let len = bytes.len(); + bytes + .try_into() + .map_err(|_| anyhow!("database encryption key must be 32 bytes, got {len}")) +} + +pub fn encrypt(key: &[u8; 32], table: &str, column: &str, id: Uuid, plain: &str) -> Result { + let mut nonce = [0; 12]; + OsRng.fill_bytes(&mut nonce); + let aad = format!("{table}:{column}:{id}"); + let cipher = Aes256Gcm::new_from_slice(key).map_err(|_| anyhow!("invalid key"))?; + let ciphertext = cipher + .encrypt( + &Nonce::from(nonce), + Payload { + msg: plain.as_bytes(), + aad: aad.as_bytes(), + }, + ) + .map_err(|_| anyhow!("encryption failed"))?; + Ok(json!({MARKER:true,"version":1,"alg":"AES-256-GCM","key_id":"s3-v1","nonce":BASE64.encode(nonce),"ciphertext":BASE64.encode(ciphertext)}).to_string()) +} + +pub fn decrypt( + key: &[u8; 32], + table: &str, + column: &str, + id: Uuid, + encoded: &str, +) -> Result { + let value: Value = serde_json::from_str(encoded)?; + ensure!(value[MARKER] == true, "missing encryption marker"); + ensure!(value["version"] == 1, "unsupported envelope version"); + ensure!( + value["alg"] == "AES-256-GCM", + "unsupported envelope algorithm" + ); + let nonce: [u8; 12] = BASE64 + .decode( + value["nonce"] + .as_str() + .ok_or_else(|| anyhow!("missing nonce"))?, + )? + .try_into() + .map_err(|_| anyhow!("invalid nonce length"))?; + let ciphertext = BASE64.decode( + value["ciphertext"] + .as_str() + .ok_or_else(|| anyhow!("missing ciphertext"))?, + )?; + let aad = format!("{table}:{column}:{id}"); + let cipher = Aes256Gcm::new_from_slice(key).map_err(|_| anyhow!("invalid key"))?; + let plaintext = cipher + .decrypt( + &Nonce::from(nonce), + Payload { + msg: &ciphertext, + aad: aad.as_bytes(), + }, + ) + .map_err(|_| anyhow!("envelope authentication failed"))?; + Ok(String::from_utf8(plaintext)?) +} + +pub fn decrypt_if_encrypted( + key: &[u8; 32], + table: &str, + column: &str, + id: Uuid, + value: String, +) -> Result { + match serde_json::from_str::(&value) { + Ok(envelope) if is_envelope(&envelope) => decrypt(key, table, column, id, &value), + _ => Ok(value), + } +} + +pub fn encrypt_json( + key: &[u8; 32], + table: &str, + column: &str, + id: Uuid, + value: &Value, +) -> Result { + Ok(serde_json::from_str(&encrypt( + key, + table, + column, + id, + &serde_json::to_string(value)?, + )?)?) +} + +pub fn decrypt_json_if_encrypted( + key: &[u8; 32], + table: &str, + column: &str, + id: Uuid, + value: Value, +) -> Result { + if !is_envelope(&value) { + return Ok(value); + } + Ok(serde_json::from_str(&decrypt( + key, + table, + column, + id, + &serde_json::to_string(&value)?, + )?)?) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn envelope_round_trips_and_authenticates_context() { + let key = [7; 32]; + let id = Uuid::new_v4(); + let encrypted = encrypt(&key, "files", "filename", id, "private.txt").unwrap(); + + assert!(!encrypted.contains("private.txt")); + assert_eq!( + decrypt(&key, "files", "filename", id, &encrypted).unwrap(), + "private.txt" + ); + assert!(decrypt(&key, "files", "storage_key", id, &encrypted).is_err()); + } + + #[test] + fn plaintext_remains_readable_during_rollout() { + assert_eq!( + decrypt_if_encrypted( + &[7; 32], + "files", + "filename", + Uuid::nil(), + "legacy.txt".into() + ) + .unwrap(), + "legacy.txt" + ); + } + + #[test] + fn user_json_with_marker_is_not_treated_as_an_envelope() { + let id = Uuid::new_v4(); + let text = json!({MARKER: true, "user": "value"}).to_string(); + assert_eq!( + decrypt_if_encrypted(&[7; 32], "files", "filename", id, text.clone()).unwrap(), + text + ); + + let json = json!({MARKER: true, "user": "value"}); + assert_eq!( + decrypt_json_if_encrypted(&[7; 32], "responses", "metadata", id, json.clone()).unwrap(), + json + ); + } + + #[test] + fn json_envelope_round_trips_and_legacy_json_remains_readable() { + let key = [9; 32]; + let id = Uuid::new_v4(); + let value = json!({"token": "secret", "nested": [1, 2, 3]}); + + let encrypted = encrypt_json(&key, "mcp_connectors", "auth_config", id, &value).unwrap(); + assert_eq!(encrypted[MARKER], true); + assert!(!encrypted.to_string().contains("secret")); + assert_eq!( + decrypt_json_if_encrypted(&key, "mcp_connectors", "auth_config", id, encrypted) + .unwrap(), + value + ); + assert_eq!( + decrypt_json_if_encrypted(&key, "mcp_connectors", "auth_config", id, value.clone()) + .unwrap(), + value + ); + } +} diff --git a/crates/database/src/lib.rs b/crates/database/src/lib.rs index 833dfefa9..67145c4b0 100644 --- a/crates/database/src/lib.rs +++ b/crates/database/src/lib.rs @@ -1,5 +1,6 @@ pub mod cluster_manager; pub mod constants; +pub mod field_encryption; pub mod migrations; pub mod mock; pub mod models; diff --git a/crates/database/src/migrations/sql/V0075__database_encryption_jobs.sql b/crates/database/src/migrations/sql/V0075__database_encryption_jobs.sql new file mode 100644 index 000000000..27b09e68f --- /dev/null +++ b/crates/database/src/migrations/sql/V0075__database_encryption_jobs.sql @@ -0,0 +1,49 @@ +CREATE TABLE database_encryption_jobs ( + id UUID PRIMARY KEY, + mode TEXT NOT NULL CHECK (mode IN ('dry_run', 'execute')), + status TEXT NOT NULL CHECK (status IN ('queued', 'running', 'completed', 'failed', 'cancelled')), + scope JSONB NOT NULL, + actions JSONB NOT NULL, + batch_size BIGINT NOT NULL CHECK (batch_size BETWEEN 1 AND 1000), + max_rows BIGINT, + cursor JSONB NOT NULL DEFAULT '{}'::jsonb, + progress JSONB NOT NULL DEFAULT '{}'::jsonb, + last_error_class TEXT, + last_error_message TEXT, + admin_actor UUID REFERENCES users(id), + operator TEXT NOT NULL DEFAULT 'admin-api', + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + started_at TIMESTAMPTZ, + completed_at TIMESTAMPTZ, + cancel_requested_at TIMESTAMPTZ +); +CREATE INDEX idx_database_encryption_jobs_status ON database_encryption_jobs(status, created_at); +CREATE UNIQUE INDEX idx_database_encryption_jobs_active_scope ON database_encryption_jobs ((scope::text)) WHERE status IN ('queued', 'running'); + +-- Encrypted envelopes are larger than the original catalog VARCHAR limits. +-- Encrypted columns are widened before repository writes or backfill jobs can +-- persist authenticated envelopes. +ALTER TABLE files + ALTER COLUMN filename TYPE TEXT, + ALTER COLUMN storage_key TYPE TEXT, + ALTER COLUMN content_type TYPE TEXT; + +ALTER TABLE mcp_connectors + ALTER COLUMN name TYPE TEXT, + ALTER COLUMN description TYPE TEXT, + ALTER COLUMN mcp_server_url TYPE TEXT, + ALTER COLUMN error_message TYPE TEXT; + +-- Keep the structural root marker queryable after response metadata is +-- encrypted. Backfill before replacing the metadata-based partial index. +ALTER TABLE responses + ADD COLUMN is_root_response BOOLEAN NOT NULL DEFAULT FALSE; + +UPDATE responses +SET is_root_response = TRUE +WHERE metadata->>'root_response' = 'true'; + +DROP INDEX IF EXISTS idx_responses_root_response_unique_per_conversation; +CREATE UNIQUE INDEX idx_responses_root_response_unique_per_conversation + ON responses(conversation_id) + WHERE is_root_response; diff --git a/crates/database/src/pool.rs b/crates/database/src/pool.rs index 2c0c8add3..0f41eae67 100644 --- a/crates/database/src/pool.rs +++ b/crates/database/src/pool.rs @@ -107,6 +107,7 @@ pub fn create_pool_with_native_tls( #[derive(Clone)] pub struct DbPool { inner: std::sync::Arc>>, + encryption_key: std::sync::Arc>>, } impl DbPool { @@ -114,6 +115,7 @@ impl DbPool { pub fn new(pool: Pool) -> Self { Self { inner: std::sync::Arc::new(std::sync::RwLock::new(Some(pool))), + encryption_key: std::sync::Arc::new(std::sync::RwLock::new(None)), } } @@ -122,6 +124,7 @@ impl DbPool { pub fn uninitialized() -> Self { Self { inner: std::sync::Arc::new(std::sync::RwLock::new(None)), + encryption_key: std::sync::Arc::new(std::sync::RwLock::new(None)), } } @@ -139,6 +142,20 @@ impl DbPool { self.inner.read().unwrap_or_else(|e| e.into_inner()).clone() } + pub fn set_encryption_key(&self, key: [u8; 32]) { + *self + .encryption_key + .write() + .unwrap_or_else(|e| e.into_inner()) = Some(key); + } + + pub fn encryption_key(&self) -> Option<[u8; 32]> { + *self + .encryption_key + .read() + .unwrap_or_else(|e| e.into_inner()) + } + /// Acquire a connection from the currently installed pool. pub async fn get(&self) -> Result { // Clone the pool out of the lock so it is not held across the await. diff --git a/crates/database/src/repositories/conversation.rs b/crates/database/src/repositories/conversation.rs index 05c38da0b..0826dbaf6 100644 --- a/crates/database/src/repositories/conversation.rs +++ b/crates/database/src/repositories/conversation.rs @@ -20,6 +20,36 @@ impl PgConversationRepository { Self { pool } } + fn encrypt_metadata( + &self, + id: Uuid, + metadata: &serde_json::Value, + ) -> Result { + match self.pool.encryption_key() { + Some(key) => crate::field_encryption::encrypt_json( + &key, + "conversations", + "metadata", + id, + metadata, + ), + None => Ok(metadata.clone()), + } + } + + fn decrypt_metadata(&self, id: Uuid, metadata: serde_json::Value) -> Result { + match self.pool.encryption_key() { + Some(key) => crate::field_encryption::decrypt_json_if_encrypted( + &key, + "conversations", + "metadata", + id, + metadata, + ), + None => Ok(metadata), + } + } + // Helper method to convert database row to Conversation model fn row_to_conversation(&self, row: tokio_postgres::Row) -> Result { let id: Uuid = row.try_get("id")?; @@ -36,7 +66,7 @@ impl PgConversationRepository { deleted_at: row.try_get("deleted_at")?, cloned_from_id: cloned_from_id.map(|id| id.into()), root_response_id: None, - metadata: row.try_get("metadata")?, + metadata: self.decrypt_metadata(id, row.try_get("metadata")?)?, created_at: row.try_get("created_at")?, updated_at: row.try_get("updated_at")?, }) @@ -53,6 +83,7 @@ impl ConversationRepository for PgConversationRepository { metadata: serde_json::Value, ) -> Result { let id = Uuid::new_v4(); + let stored_metadata = self.encrypt_metadata(id, &metadata)?; let row = retry_db!("create_new_conversation", { let now = Utc::now(); @@ -70,7 +101,7 @@ impl ConversationRepository for PgConversationRepository { VALUES ($1, $2, $3, $4, $5, $6) RETURNING * "#, - &[&id, &workspace_id.0, &api_key_id, &metadata, &now, &now], + &[&id, &workspace_id.0, &api_key_id, &stored_metadata, &now, &now], ) .await .map_err(map_db_error) @@ -119,6 +150,7 @@ impl ConversationRepository for PgConversationRepository { workspace_id: WorkspaceId, metadata: serde_json::Value, ) -> Result> { + let stored_metadata = self.encrypt_metadata(id.0, &metadata)?; let row = retry_db!("update_conversation_metadata", { let now = Utc::now(); let client = self @@ -136,7 +168,7 @@ impl ConversationRepository for PgConversationRepository { WHERE id = $1 AND workspace_id = $2 AND deleted_at IS NULL RETURNING * "#, - &[&id.0, &workspace_id.0, &metadata, &now], + &[&id.0, &workspace_id.0, &stored_metadata, &now], ) .await .map_err(map_db_error) @@ -266,47 +298,54 @@ impl ConversationRepository for PgConversationRepository { .await .context("Failed to start transaction")?; - // Step 1: Clone the conversation with a new ID and append " (Copy)" to title in metadata - // Reset pinned_at, archived_at, deleted_at to NULL for the clone - let conv_row = transaction + // Decrypt and re-encrypt metadata because the envelope AAD includes the row ID. + let original = transaction .query_opt( - r#" - INSERT INTO conversations (id, workspace_id, api_key_id, pinned_at, archived_at, deleted_at, cloned_from_id, metadata, created_at, updated_at) - SELECT - $1, - workspace_id, - $2, - NULL, - NULL, - NULL, - id, - CASE - WHEN metadata->>'title' IS NOT NULL THEN - jsonb_set(metadata, '{title}', to_jsonb((metadata->>'title') || ' (Copy)')) - ELSE - metadata - END, - $3, - $4 - FROM conversations - WHERE id = $5 AND workspace_id = $6 AND deleted_at IS NULL - RETURNING * - "#, - &[&new_conv_id, &api_key_id, &now, &now, &id.0, &workspace_id.0], + "SELECT metadata FROM conversations WHERE id = $1 AND workspace_id = $2 AND deleted_at IS NULL", + &[&id.0, &workspace_id.0], ) .await - .context("Failed to clone conversation")?; - - if conv_row.is_none() { + .context("Failed to load conversation metadata for clone")?; + let Some(original) = original else { // Conversation not found or is deleted, rollback and return None transaction.rollback().await.ok(); return Ok(None); + }; + let mut metadata = self.decrypt_metadata(id.0, original.try_get("metadata")?)?; + if let Some(title) = metadata.get_mut("title") { + if let Some(title) = title.as_str() { + *metadata.get_mut("title").expect("title exists") = + serde_json::Value::String(format!("{title} (Copy)")); + } } + let stored_metadata = self.encrypt_metadata(new_conv_id, &metadata)?; + transaction + .execute( + r#" + INSERT INTO conversations ( + id, workspace_id, api_key_id, pinned_at, archived_at, deleted_at, + cloned_from_id, metadata, created_at, updated_at + ) + SELECT $1, workspace_id, $2, NULL, NULL, NULL, id, $3, $4, $4 + FROM conversations + WHERE id = $5 AND workspace_id = $6 AND deleted_at IS NULL + "#, + &[ + &new_conv_id, + &api_key_id, + &stored_metadata, + &now, + &id.0, + &workspace_id.0, + ], + ) + .await + .context("Failed to clone conversation")?; // Step 2: Get all responses from the original conversation let original_responses = transaction .query( - "SELECT id FROM responses WHERE conversation_id = $1 AND workspace_id = $2 ORDER BY created_at ASC", + "SELECT id, instructions, metadata FROM responses WHERE conversation_id = $1 AND workspace_id = $2 ORDER BY created_at ASC", &[&id.0, &workspace_id.0], ) .await @@ -319,30 +358,83 @@ impl ConversationRepository for PgConversationRepository { let old_response_id: Uuid = orig_row.try_get("id")?; let new_response_id = Uuid::new_v4(); id_map.insert(old_response_id, new_response_id); + let instructions: Option = orig_row.try_get("instructions")?; + let metadata: serde_json::Value = orig_row.try_get("metadata")?; + let (stored_instructions, stored_response_metadata) = + if let Some(key) = self.pool.encryption_key() { + let instructions = instructions + .map(|value| { + crate::field_encryption::decrypt_if_encrypted( + &key, + "responses", + "instructions", + old_response_id, + value, + ) + .and_then(|plain| { + crate::field_encryption::encrypt( + &key, + "responses", + "instructions", + new_response_id, + &plain, + ) + }) + }) + .transpose()?; + let metadata = crate::field_encryption::decrypt_json_if_encrypted( + &key, + "responses", + "metadata", + old_response_id, + metadata, + )?; + ( + instructions, + crate::field_encryption::encrypt_json( + &key, + "responses", + "metadata", + new_response_id, + &metadata, + )?, + ) + } else { + (instructions, metadata) + }; // Clone the response with new ID and new conversation_id transaction .execute( r#" - INSERT INTO responses (id, workspace_id, api_key_id, model, status, instructions, conversation_id, previous_response_id, next_response_ids, usage, metadata, created_at, updated_at) + INSERT INTO responses (id, workspace_id, api_key_id, model, status, instructions, conversation_id, previous_response_id, next_response_ids, usage, metadata, is_root_response, created_at, updated_at) SELECT $1, workspace_id, $2, model, status, - instructions, + $4, $3, previous_response_id, next_response_ids, usage, - metadata, - $4, - $5 + $5, + is_root_response, + $6, + $6 FROM responses - WHERE id = $6 + WHERE id = $7 "#, - &[&new_response_id, &api_key_id, &new_conv_id, &now, &now, &old_response_id], + &[ + &new_response_id, + &api_key_id, + &new_conv_id, + &stored_instructions, + &stored_response_metadata, + &now, + &old_response_id, + ], ) .await .context("Failed to clone response")?; @@ -410,6 +502,7 @@ impl ConversationRepository for PgConversationRepository { .context("Failed to get original response items")?; for item_row in &original_items { + let old_item_id: Uuid = item_row.try_get("id")?; let old_response_id: Uuid = item_row.try_get("response_id")?; let mut item_json: serde_json::Value = item_row.try_get("item")?; let original_created_at: chrono::DateTime = item_row.try_get("created_at")?; @@ -421,6 +514,16 @@ impl ConversationRepository for PgConversationRepository { .unwrap_or(old_response_id); let new_item_id = Uuid::new_v4(); + if let Some(key) = self.pool.encryption_key() { + item_json = crate::field_encryption::decrypt_json_if_encrypted( + &key, + "response_items", + "item", + old_item_id, + item_json, + )?; + } + // Update the "id" field inside the item JSON to use the new item ID // The item JSON has a structure like: { "id": "msg_...", "type": "message", ... } if let Some(obj) = item_json.as_object_mut() { @@ -429,6 +532,16 @@ impl ConversationRepository for PgConversationRepository { obj.insert("id".to_string(), serde_json::Value::String(new_msg_id)); } + if let Some(key) = self.pool.encryption_key() { + item_json = crate::field_encryption::encrypt_json( + &key, + "response_items", + "item", + new_item_id, + &item_json, + )?; + } + transaction .execute( r#" diff --git a/crates/database/src/repositories/file.rs b/crates/database/src/repositories/file.rs index a172a426f..3e481d27d 100644 --- a/crates/database/src/repositories/file.rs +++ b/crates/database/src/repositories/file.rs @@ -23,6 +23,34 @@ impl FileRepository { params: services::files::ports::CreateFileParams, ) -> Result { let id = Uuid::new_v4(); + let (filename, content_type, storage_key) = if let Some(key) = self.pool.encryption_key() { + ( + crate::field_encryption::encrypt(&key, "files", "filename", id, ¶ms.filename) + .map_err(RepositoryError::DataConversionError)?, + crate::field_encryption::encrypt( + &key, + "files", + "content_type", + id, + ¶ms.content_type, + ) + .map_err(RepositoryError::DataConversionError)?, + crate::field_encryption::encrypt( + &key, + "files", + "storage_key", + id, + ¶ms.storage_key, + ) + .map_err(RepositoryError::DataConversionError)?, + ) + } else { + ( + params.filename.clone(), + params.content_type.clone(), + params.storage_key.clone(), + ) + }; let row = match retry_db!("create_new_file_record", { let now = Utc::now(); @@ -45,11 +73,11 @@ impl FileRepository { "#, &[ &id, - ¶ms.filename, + &filename, ¶ms.bytes, - ¶ms.content_type, + &content_type, ¶ms.purpose, - ¶ms.storage_key, + &storage_key, ¶ms.workspace_id, ¶ms.uploaded_by_api_key_id, &now, @@ -359,13 +387,36 @@ impl FileRepository { /// Helper function to convert database row to File fn row_to_file(&self, row: tokio_postgres::Row) -> Result { + let id: Uuid = row.get("id"); + let mut filename: String = row.get("filename"); + let mut content_type: String = row.get("content_type"); + let mut storage_key: String = row.get("storage_key"); + if let Some(key) = self.pool.encryption_key() { + filename = crate::field_encryption::decrypt_if_encrypted( + &key, "files", "filename", id, filename, + )?; + content_type = crate::field_encryption::decrypt_if_encrypted( + &key, + "files", + "content_type", + id, + content_type, + )?; + storage_key = crate::field_encryption::decrypt_if_encrypted( + &key, + "files", + "storage_key", + id, + storage_key, + )?; + } Ok(File { - id: row.get("id"), - filename: row.get("filename"), + id, + filename, bytes: row.get("bytes"), - content_type: row.get("content_type"), + content_type, purpose: row.get("purpose"), - storage_key: row.get("storage_key"), + storage_key, workspace_id: row.get("workspace_id"), uploaded_by_api_key_id: row.get("uploaded_by_api_key_id"), created_at: row.get("created_at"), diff --git a/crates/database/src/repositories/mcp_connector.rs b/crates/database/src/repositories/mcp_connector.rs index e9908c031..3d023645d 100644 --- a/crates/database/src/repositories/mcp_connector.rs +++ b/crates/database/src/repositories/mcp_connector.rs @@ -20,6 +20,50 @@ impl McpConnectorRepository { Self { pool } } + fn encrypt_text(&self, table: &str, column: &str, id: Uuid, value: String) -> Result { + match self.pool.encryption_key() { + Some(key) => crate::field_encryption::encrypt(&key, table, column, id, &value), + None => Ok(value), + } + } + + fn decrypt_text(&self, table: &str, column: &str, id: Uuid, value: String) -> Result { + match self.pool.encryption_key() { + Some(key) => { + crate::field_encryption::decrypt_if_encrypted(&key, table, column, id, value) + } + None => Ok(value), + } + } + + fn encrypt_json( + &self, + table: &str, + column: &str, + id: Uuid, + value: serde_json::Value, + ) -> Result { + match self.pool.encryption_key() { + Some(key) => crate::field_encryption::encrypt_json(&key, table, column, id, &value), + None => Ok(value), + } + } + + fn decrypt_json( + &self, + table: &str, + column: &str, + id: Uuid, + value: serde_json::Value, + ) -> Result { + match self.pool.encryption_key() { + Some(key) => { + crate::field_encryption::decrypt_json_if_encrypted(&key, table, column, id, value) + } + None => Ok(value), + } + } + /// Create a new MCP connector for an organization pub async fn create( &self, @@ -30,12 +74,25 @@ impl McpConnectorRepository { let id = Uuid::new_v4(); // Convert bearer token to auth_config if present - let auth_config = request.bearer_token.as_ref().map(|token| { + let mut auth_config = request.bearer_token.as_ref().map(|token| { serde_json::to_value(McpBearerConfig { token: token.clone(), }) .unwrap() }); + let description = request + .description + .map(|value| self.encrypt_text("mcp_connectors", "description", id, value)) + .transpose()?; + let mcp_server_url = self.encrypt_text( + "mcp_connectors", + "mcp_server_url", + id, + request.mcp_server_url, + )?; + auth_config = auth_config + .map(|value| self.encrypt_json("mcp_connectors", "auth_config", id, value)) + .transpose()?; let row = retry_db!("create_mcp_connector", { let now = Utc::now(); @@ -62,8 +119,8 @@ impl McpConnectorRepository { &id, &organization_id, &request.name, - &request.description, - &request.mcp_server_url, + &description, + &mcp_server_url, &request.auth_type.to_string(), &auth_config, &creator_user_id, @@ -172,6 +229,20 @@ impl McpConnectorRepository { request: UpdateMcpConnectorRequest, ) -> Result { let now = Utc::now(); + let encrypted_description = request + .description + .map(|value| self.encrypt_text("mcp_connectors", "description", id, value)) + .transpose()?; + let encrypted_url = request + .mcp_server_url + .map(|value| self.encrypt_text("mcp_connectors", "mcp_server_url", id, value)) + .transpose()?; + let encrypted_auth_config = request + .bearer_token + .map(|token| serde_json::to_value(McpBearerConfig { token })) + .transpose()? + .map(|value| self.encrypt_json("mcp_connectors", "auth_config", id, value)) + .transpose()?; // Build dynamic update query let mut query = String::from("UPDATE mcp_connectors SET updated_at = $1"); @@ -184,13 +255,13 @@ impl McpConnectorRepository { param_idx += 1; } - if let Some(ref description) = request.description { + if let Some(ref description) = encrypted_description { query.push_str(&format!(", description = ${param_idx}")); params.push(description); param_idx += 1; } - if let Some(ref url) = request.mcp_server_url { + if let Some(ref url) = encrypted_url { query.push_str(&format!(", mcp_server_url = ${param_idx}")); params.push(url); param_idx += 1; @@ -204,14 +275,9 @@ impl McpConnectorRepository { param_idx += 1; } - let auth_config; - if let Some(ref bearer_token) = request.bearer_token { - auth_config = serde_json::to_value(McpBearerConfig { - token: bearer_token.clone(), - }) - .unwrap(); + if let Some(ref auth_config) = encrypted_auth_config { query.push_str(&format!(", auth_config = ${param_idx}")); - params.push(&auth_config); + params.push(auth_config); param_idx += 1; } @@ -293,6 +359,12 @@ impl McpConnectorRepository { // Use separate variables to avoid type ambiguity let update_last_connected = status == McpConnectionStatus::Connected; + let error_message = error_message + .map(|value| self.encrypt_text("mcp_connectors", "error_message", id, value)) + .transpose()?; + let capabilities = capabilities + .map(|value| self.encrypt_json("mcp_connectors", "capabilities", id, value)) + .transpose()?; let rows_affected = retry_db!("update_mcp_connector_connection_status", { let now = Utc::now(); @@ -375,6 +447,15 @@ impl McpConnectorRepository { duration_ms: Option, ) -> Result<()> { let id = Uuid::new_v4(); + let request_payload = request_payload + .map(|value| self.encrypt_json("mcp_connector_usage", "request_payload", id, value)) + .transpose()?; + let response_payload = response_payload + .map(|value| self.encrypt_json("mcp_connector_usage", "response_payload", id, value)) + .transpose()?; + let error_message = error_message + .map(|value| self.encrypt_text("mcp_connector_usage", "error_message", id, value)) + .transpose()?; retry_db!("log_mcp_connector_usage", { let now = Utc::now(); @@ -449,6 +530,7 @@ impl McpConnectorRepository { /// Convert a database row to McpConnector fn row_to_connector(&self, row: tokio_postgres::Row) -> Result { + let id: Uuid = row.get("id"); let auth_type_str: String = row.get("auth_type"); let auth_type = match auth_type_str.as_str() { "none" => McpAuthType::None, @@ -465,36 +547,68 @@ impl McpConnectorRepository { }; Ok(McpConnector { - id: row.get("id"), + id, organization_id: row.get("organization_id"), name: row.get("name"), - description: row.get("description"), - mcp_server_url: row.get("mcp_server_url"), + description: row + .get::<_, Option>("description") + .map(|value| self.decrypt_text("mcp_connectors", "description", id, value)) + .transpose()?, + mcp_server_url: self.decrypt_text( + "mcp_connectors", + "mcp_server_url", + id, + row.get("mcp_server_url"), + )?, auth_type, - auth_config: row.get("auth_config"), + auth_config: row + .get::<_, Option>("auth_config") + .map(|value| self.decrypt_json("mcp_connectors", "auth_config", id, value)) + .transpose()?, is_active: row.get("is_active"), created_by: row.get("created_by"), created_at: row.get("created_at"), updated_at: row.get("updated_at"), last_connected_at: row.get("last_connected_at"), connection_status, - error_message: row.get("error_message"), - capabilities: row.get("capabilities"), - metadata: row.get("metadata"), + error_message: row + .get::<_, Option>("error_message") + .map(|value| self.decrypt_text("mcp_connectors", "error_message", id, value)) + .transpose()?, + capabilities: row + .get::<_, Option>("capabilities") + .map(|value| self.decrypt_json("mcp_connectors", "capabilities", id, value)) + .transpose()?, + metadata: row + .get::<_, Option>("metadata") + .map(|value| self.decrypt_json("mcp_connectors", "metadata", id, value)) + .transpose()?, }) } /// Convert a database row to McpConnectorUsage fn row_to_usage(&self, row: tokio_postgres::Row) -> Result { + let id: Uuid = row.get("id"); Ok(McpConnectorUsage { - id: row.get("id"), + id, connector_id: row.get("connector_id"), user_id: row.get("user_id"), method: row.get("method"), - request_payload: row.get("request_payload"), - response_payload: row.get("response_payload"), + request_payload: row + .get::<_, Option>("request_payload") + .map(|value| self.decrypt_json("mcp_connector_usage", "request_payload", id, value)) + .transpose()?, + response_payload: row + .get::<_, Option>("response_payload") + .map(|value| { + self.decrypt_json("mcp_connector_usage", "response_payload", id, value) + }) + .transpose()?, status_code: row.get("status_code"), - error_message: row.get("error_message"), + error_message: row + .get::<_, Option>("error_message") + .map(|value| self.decrypt_text("mcp_connector_usage", "error_message", id, value)) + .transpose()?, duration_ms: row.get("duration_ms"), created_at: row.get("created_at"), }) diff --git a/crates/database/src/repositories/response.rs b/crates/database/src/repositories/response.rs index 07517ea1a..f916938d8 100644 --- a/crates/database/src/repositories/response.rs +++ b/crates/database/src/repositories/response.rs @@ -19,8 +19,65 @@ impl PgResponseRepository { Self { pool } } + fn encrypt_text(&self, id: Uuid, column: &str, value: Option<&str>) -> Result> { + value + .map(|value| match self.pool.encryption_key() { + Some(key) => crate::field_encryption::encrypt(&key, "responses", column, id, value), + None => Ok(value.to_string()), + }) + .transpose() + } + + fn decrypt_text( + &self, + id: Uuid, + column: &str, + value: Option, + ) -> Result> { + value + .map(|value| match self.pool.encryption_key() { + Some(key) => crate::field_encryption::decrypt_if_encrypted( + &key, + "responses", + column, + id, + value, + ), + None => Ok(value), + }) + .transpose() + } + + fn encrypt_metadata(&self, id: Uuid, value: &serde_json::Value) -> Result { + match self.pool.encryption_key() { + Some(key) => { + crate::field_encryption::encrypt_json(&key, "responses", "metadata", id, value) + } + None => Ok(value.clone()), + } + } + + fn decrypt_metadata( + &self, + id: Uuid, + value: Option, + ) -> Result> { + value + .map(|value| match self.pool.encryption_key() { + Some(key) => crate::field_encryption::decrypt_json_if_encrypted( + &key, + "responses", + "metadata", + id, + value, + ), + None => Ok(value), + }) + .transpose() + } + /// Fetch the ID of the structural root response for a conversation, if it exists. - /// Root rows are identified by metadata->>'root_response' = 'true'. + /// Root rows are identified by the dedicated structural flag. async fn fetch_root_id_opt( &self, conversation_uuid: Uuid, @@ -41,7 +98,7 @@ impl PgResponseRepository { FROM responses WHERE conversation_id = $1 AND workspace_id = $2 - AND metadata->>'root_response' = 'true' + AND is_root_response ORDER BY created_at ASC LIMIT 1 "#, @@ -85,6 +142,10 @@ impl PgResponseRepository { let metadata_json = serde_json::json!({ "root_response": true }); + let root_id = Uuid::new_v4(); + let stored_metadata = self + .encrypt_metadata(root_id, &metadata_json) + .map_err(RepositoryError::DataConversionError)?; let next_response_ids_json = serde_json::json!([]); let inserted_row_opt = retry_db!("insert_conversation_root", { @@ -102,16 +163,17 @@ impl PgResponseRepository { .query_opt( r#" INSERT INTO responses ( - workspace_id, api_key_id, model, status, instructions, conversation_id, - previous_response_id, next_response_ids, usage, metadata, + id, workspace_id, api_key_id, model, status, instructions, conversation_id, + previous_response_id, next_response_ids, usage, metadata, is_root_response, created_at, updated_at ) - VALUES ($1, $2, $3, $4, NULL, $5, NULL, $6, $7, $8, $9, $9) - ON CONFLICT (conversation_id) WHERE metadata->>'root_response' = 'true' + VALUES ($1, $2, $3, $4, $5, NULL, $6, NULL, $7, $8, $9, TRUE, $10, $10) + ON CONFLICT (conversation_id) WHERE is_root_response DO NOTHING RETURNING id "#, &[ + &root_id, &workspace_id.0, api_key_id, &model, @@ -119,7 +181,7 @@ impl PgResponseRepository { &conversation_uuid, &next_response_ids_json, &usage_json, - &metadata_json, + &stored_metadata, &now, ], ) @@ -300,6 +362,12 @@ impl ResponseRepositoryTrait for PgResponseRepository { "total_tokens": 0 }); let metadata_json = request.metadata.unwrap_or_else(|| serde_json::json!({})); + let stored_instructions = self.encrypt_text( + response_uuid, + "instructions", + request.instructions.as_deref(), + )?; + let stored_metadata = self.encrypt_metadata(response_uuid, &metadata_json)?; let next_response_ids_json = serde_json::json!([]); // Insert response and update previous response in a single retry_db! block @@ -330,12 +398,12 @@ impl ResponseRepositoryTrait for PgResponseRepository { &api_key_id, &request.model, &status, - &request.instructions, + &stored_instructions, &conversation_uuid, &previous_response_uuid, &next_response_ids_json, &usage_json, - &metadata_json, + &stored_metadata, &now, &now, ], @@ -478,9 +546,10 @@ impl ResponseRepositoryTrait for PgResponseRepository { let created_at: DateTime = row.get("created_at"); let conversation_uuid: Option = row.get("conversation_id"); let usage_json: Option = row.get("usage"); - let metadata_json: Option = row.get("metadata"); + let metadata_json = self.decrypt_metadata(response_uuid, row.get("metadata"))?; let model: String = row.get("model"); - let instructions: Option = row.get("instructions"); + let instructions = + self.decrypt_text(response_uuid, "instructions", row.get("instructions"))?; let previous_response_uuid: Option = row.get("previous_response_id"); let next_response_ids_json: Option = row.get("next_response_ids"); @@ -658,8 +727,8 @@ impl ResponseRepositoryTrait for PgResponseRepository { }; // Parse metadata - let metadata_value: Option = row.get(10); - let metadata = metadata_value; + let metadata = self.decrypt_metadata(response_uuid, row.get(10))?; + let instructions = self.decrypt_text(response_uuid, "instructions", row.get(5))?; // Parse status let status_str: String = row.get(4); @@ -680,7 +749,7 @@ impl ResponseRepositoryTrait for PgResponseRepository { conversation: conversation_ref, error: None, incomplete_details: None, - instructions: row.get(5), + instructions, max_output_tokens: None, // Not stored in DB max_tool_calls: None, // Not stored in DB model: row.get(3), @@ -784,7 +853,7 @@ impl ResponseRepositoryTrait for PgResponseRepository { FROM responses WHERE conversation_id = $1 AND workspace_id = $2 - AND COALESCE((metadata->>'root_response')::boolean, false) = false + AND NOT is_root_response ORDER BY created_at DESC LIMIT 1 "#, @@ -812,9 +881,10 @@ impl ResponseRepositoryTrait for PgResponseRepository { let created_at: DateTime = row.get("created_at"); let conversation_uuid: Option = row.get("conversation_id"); let usage_json: Option = row.get("usage"); - let metadata_json: Option = row.get("metadata"); + let metadata_json = self.decrypt_metadata(response_uuid, row.get("metadata"))?; let model: String = row.get("model"); - let instructions: Option = row.get("instructions"); + let instructions = + self.decrypt_text(response_uuid, "instructions", row.get("instructions"))?; let previous_response_uuid: Option = row.get("previous_response_id"); let next_response_ids_json: Option = row.get("next_response_ids"); diff --git a/crates/database/src/repositories/response_item.rs b/crates/database/src/repositories/response_item.rs index e58ad72ed..e5c4b76d0 100644 --- a/crates/database/src/repositories/response_item.rs +++ b/crates/database/src/repositories/response_item.rs @@ -35,6 +35,7 @@ //! let items = repo.list_by_response(response_id).await?; //! ``` +use crate::field_encryption; use crate::pool::DbPool; use crate::repositories::utils::map_db_error; use crate::retry_db; @@ -60,7 +61,17 @@ impl PgResponseItemsRepository { /// Helper method to convert database row to ResponseOutputItem /// Enriches the item with response metadata (response_id, previous_response_id, next_response_ids, created_at) fn row_to_item(&self, row: tokio_postgres::Row) -> Result { - let item_json: serde_json::Value = row.try_get("item")?; + let id: Uuid = row.try_get("id")?; + let mut item_json: serde_json::Value = row.try_get("item")?; + if let Some(key) = self.pool.encryption_key() { + item_json = field_encryption::decrypt_json_if_encrypted( + &key, + "response_items", + "item", + id, + item_json, + )?; + } let mut item: ResponseOutputItem = serde_json::from_value(item_json) .context("Failed to deserialize response item from database")?; @@ -256,7 +267,12 @@ impl ResponseItemRepositoryTrait for PgResponseItemsRepository { let id = Self::extract_uuid_from_item_id(item_id); // Serialize the item to JSON for storage - let item_json = serde_json::to_value(&item).context("Failed to serialize response item")?; + let mut item_json = + serde_json::to_value(&item).context("Failed to serialize response item")?; + if let Some(key) = self.pool.encryption_key() { + item_json = + field_encryption::encrypt_json(&key, "response_items", "item", id, &item_json)?; + } let conversation_uuid = conversation_id.map(|cid| cid.0); @@ -353,7 +369,12 @@ impl ResponseItemRepositoryTrait for PgResponseItemsRepository { item: ResponseOutputItem, ) -> Result { // Serialize the updated item to JSON - let item_json = serde_json::to_value(&item).context("Failed to serialize response item")?; + let mut item_json = + serde_json::to_value(&item).context("Failed to serialize response item")?; + if let Some(key) = self.pool.encryption_key() { + item_json = + field_encryption::encrypt_json(&key, "response_items", "item", id.0, &item_json)?; + } let row = retry_db!("update_response_item", { let now = Utc::now(); @@ -367,10 +388,20 @@ impl ResponseItemRepositoryTrait for PgResponseItemsRepository { client .query_opt( r#" - UPDATE response_items - SET item = $2, updated_at = $3 - WHERE id = $1 - RETURNING * + WITH updated AS ( + UPDATE response_items + SET item = $2, updated_at = $3 + WHERE id = $1 + RETURNING * + ) + SELECT + updated.*, + r.previous_response_id, + r.next_response_ids, + r.created_at AS response_created_at, + r.model + FROM updated + JOIN responses r ON updated.response_id = r.id "#, &[&id.0, &item_json, &now], )