diff --git a/bindings/python/src/lib.rs b/bindings/python/src/lib.rs index 37143fb399..cb35ef95fc 100644 --- a/bindings/python/src/lib.rs +++ b/bindings/python/src/lib.rs @@ -283,6 +283,7 @@ impl PyOracleConfig { pool_min: self.pool_min, pool_max: self.pool_max, pool_timeout_secs: self.pool_timeout_secs, + schema: None, } } } @@ -321,6 +322,7 @@ impl PyRedisConfig { url: self.url.clone(), pool_max: self.pool_max, retention_days: self.retention_days, + schema: None, } } } @@ -353,6 +355,7 @@ impl PyPostgresConfig { config::PostgresConfig { db_url: self.db_url.clone().unwrap_or_default(), pool_max: self.pool_max, + schema: None, } } } diff --git a/data_connector/README.md b/data_connector/README.md index 49f7943a2e..f4dfefb0d4 100644 --- a/data_connector/README.md +++ b/data_connector/README.md @@ -107,6 +107,7 @@ from JSON/YAML). Each database backend has a dedicated config struct. |-------|------|-------------| | `db_url` | `String` | Connection URL (`postgres://user:pass@host:port/dbname`). Validated for scheme, host, and database name. | | `pool_max` | `usize` | Maximum connections in the deadpool pool (default helper: 16). Must be > 0. | +| `schema` | `Option` | Optional schema customization. See [Schema Configuration](#schema-configuration). | Call `validate()` to check the URL before use. @@ -117,6 +118,7 @@ Call `validate()` to check the URL before use. | `url` | `String` | -- | Connection URL (`redis://` or `rediss://`). | | `pool_max` | `usize` | 16 | Maximum pool connections. | | `retention_days` | `Option` | `Some(30)` | TTL in days for stored data. `None` disables expiration. | +| `schema` | `Option` | `None` | Optional schema customization. See [Schema Configuration](#schema-configuration). | Call `validate()` to check the URL before use. @@ -132,6 +134,61 @@ Call `validate()` to check the URL before use. | `pool_min` | `usize` | 1 | Minimum pool connections. | | `pool_max` | `usize` | 16 | Maximum pool connections. | | `pool_timeout_secs` | `u64` | 30 | Connection acquisition timeout in seconds. | +| `schema` | `Option` | `None` | Optional schema customization. See [Schema Configuration](#schema-configuration). | + +### Schema Configuration + +All three database backends (Oracle, Postgres, Redis) accept an optional +`SchemaConfig` that lets you customize table names and column names without +modifying source code. When `schema` is omitted, all backends use their +default table and column names — zero behavioral change. + +```yaml +postgres: + db_url: "postgres://user:pass@localhost:5432/mydb" + pool_max: 16 + + schema: + owner: "myschema" # Oracle: schema prefix (MYSCHEMA."TABLE") + # Redis: key prefix ("myschema:conversation:{id}") + # Postgres: ignored (use search_path for schema control) + + conversations: + table: "my_conversations" # Overrides default "conversations" + columns: + id: "conv_id" # Overrides column name "id" -> "conv_id" + metadata: "conv_meta" # Overrides column name "metadata" -> "conv_meta" + + responses: + table: "my_responses" + columns: + safety_identifier: "user_identifier" + + conversation_items: + table: "my_items" + + conversation_item_links: + table: "my_links" +``` + +`SchemaConfig` has two types: + +| Type | Fields | Purpose | +|------|--------|---------| +| `SchemaConfig` | `owner`, `conversations`, `responses`, `conversation_items`, `conversation_item_links` | Top-level config with an optional owner/prefix and per-table settings | +| `TableConfig` | `table`, `columns` | Per-table config: physical table name and a map of logical-to-physical column name overrides | + +Key behaviors: + +- **`col(field)`** returns the physical column name for a logical field name. + If no override is configured, the logical name is returned unchanged. +- **`qualified_table(owner)`** returns `OWNER."TABLE"` when an owner is set + (used by Oracle), or just the table name otherwise. +- **Validation** runs at startup. All identifiers must match `[a-zA-Z0-9_]+`. + Invalid identifiers are rejected before any queries execute. +- **Redis**: Only `owner` (key prefix) and `columns` (hash field names) affect + Redis behavior. The `table` field is ignored for Redis key patterns — keys + always use hardcoded entity names (`conversation`, `item`, `response`). ## Data Model @@ -181,10 +238,10 @@ struct reconstructs the chronological sequence of related responses. ## Database Schema All database backends auto-create their schemas on first connection. The -following tables are used: +following default table names are used (configurable via `SchemaConfig`): -| Table | Purpose | -|-------|---------| +| Default Table | Purpose | +|---------------|---------| | `conversations` | Conversation records with metadata | | `conversation_items` | Individual items (messages, tool calls, etc.) | | `conversation_item_links` | Join table linking items to conversations with ordering (`added_at`) | @@ -194,6 +251,11 @@ PostgreSQL additionally creates an index on `conversation_item_links(conversation_id, added_at)` for efficient cursor-based listing. +Column names within each table can also be overridden via `SchemaConfig`. The +config describes the existing database schema — it does not perform migrations. +If you rename a column in config, the corresponding database column must +already exist with that name. + ## Testing Run the unit tests (Memory and NoOp backends, config validation, ID generation): diff --git a/data_connector/src/common.rs b/data_connector/src/common.rs index 498210e035..1638916732 100644 --- a/data_connector/src/common.rs +++ b/data_connector/src/common.rs @@ -2,7 +2,41 @@ use std::collections::HashMap; use serde_json::Value; -use crate::core::ConversationMetadata; +use crate::{core::ConversationMetadata, schema::SchemaConfig}; + +/// Logical column names for the responses table, in canonical SELECT order. +/// +/// Shared between Oracle and Postgres backends to build dynamic SELECT queries. +/// The order here doesn't affect correctness (both backends read by name, not +/// position), but having a single source prevents accidental divergence. +pub(super) const RESPONSE_COLUMNS: &[&str] = &[ + "id", + "conversation_id", + "previous_response_id", + "input", + "instructions", + "output", + "tool_calls", + "metadata", + "created_at", + "safety_identifier", + "model", + "raw_response", +]; + +/// Build the `SELECT col1, col2, ... FROM table` base query for responses. +/// +/// Used by Oracle and Postgres to pre-build the SELECT prefix at construction +/// time, avoiding repeated string formatting on every query. +pub(super) fn build_response_select_base(schema: &SchemaConfig) -> String { + let s = &schema.responses; + let table = s.qualified_table(schema.owner.as_deref()); + let cols: Vec<&str> = RESPONSE_COLUMNS + .iter() + .map(|&logical| s.col(logical)) + .collect(); + format!("SELECT {} FROM {table}", cols.join(", ")) +} /// Parse raw JSON string into `ConversationMetadata` (`JsonMap`). /// diff --git a/data_connector/src/config.rs b/data_connector/src/config.rs index 08328cc7d8..90d7741711 100644 --- a/data_connector/src/config.rs +++ b/data_connector/src/config.rs @@ -3,6 +3,8 @@ use serde::{Deserialize, Serialize}; use url::Url; +use crate::schema::SchemaConfig; + /// History backend configuration #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)] #[serde(rename_all = "lowercase")] @@ -33,6 +35,9 @@ pub struct OracleConfig { pub pool_max: usize, #[serde(default = "default_pool_timeout_secs")] pub pool_timeout_secs: u64, + /// Optional schema customization (table names, column names, extra columns). + #[serde(default, skip_serializing_if = "Option::is_none")] + pub schema: Option, } impl OracleConfig { @@ -71,6 +76,7 @@ impl std::fmt::Debug for OracleConfig { .field("pool_min", &self.pool_min) .field("pool_max", &self.pool_max) .field("pool_timeout_secs", &self.pool_timeout_secs) + .field("schema", &self.schema) .finish() } } @@ -82,6 +88,9 @@ pub struct PostgresConfig { pub db_url: String, // Database pool max size pub pool_max: usize, + /// Optional schema customization (table names, column names, extra columns). + #[serde(default, skip_serializing_if = "Option::is_none")] + pub schema: Option, } impl PostgresConfig { @@ -131,6 +140,9 @@ pub struct RedisConfig { // Connection pool max size #[serde(default = "default_redis_pool_max")] pub pool_max: usize, + /// Optional schema customization (key prefix, field names, extra fields). + #[serde(default, skip_serializing_if = "Option::is_none")] + pub schema: Option, // Data retention in days. If None, data persists indefinitely. #[serde(default = "default_redis_retention_days")] pub retention_days: Option, @@ -185,6 +197,7 @@ mod tests { let cfg = PostgresConfig { db_url: "postgres://user:pass@localhost:5432/mydb".to_string(), pool_max: 16, + schema: None, }; cfg.validate() .expect("valid postgres URL should pass validation"); @@ -195,6 +208,7 @@ mod tests { let cfg = PostgresConfig { db_url: "postgresql://user:pass@localhost/mydb".to_string(), pool_max: 8, + schema: None, }; cfg.validate() .expect("postgresql:// scheme should also be accepted"); @@ -205,6 +219,7 @@ mod tests { let cfg = PostgresConfig { db_url: " ".to_string(), pool_max: 16, + schema: None, }; let err = cfg.validate().expect_err("empty URL should fail"); assert!( @@ -218,6 +233,7 @@ mod tests { let cfg = PostgresConfig { db_url: "mysql://user:pass@localhost/mydb".to_string(), pool_max: 16, + schema: None, }; let err = cfg.validate().expect_err("mysql scheme should be rejected"); assert!( @@ -232,6 +248,7 @@ mod tests { let cfg = PostgresConfig { db_url: "postgres:///mydb".to_string(), pool_max: 16, + schema: None, }; let err = cfg.validate().expect_err("missing host should fail"); assert!( @@ -245,6 +262,7 @@ mod tests { let cfg = PostgresConfig { db_url: "postgres://user:pass@localhost".to_string(), pool_max: 16, + schema: None, }; let err = cfg .validate() @@ -260,6 +278,7 @@ mod tests { let cfg = PostgresConfig { db_url: "postgres://user:pass@localhost/mydb".to_string(), pool_max: 0, + schema: None, }; let err = cfg.validate().expect_err("pool_max=0 should fail"); assert!( @@ -276,6 +295,7 @@ mod tests { url: "redis://:password@localhost:6379/0".to_string(), pool_max: 16, retention_days: Some(30), + schema: None, }; cfg.validate() .expect("valid redis URL should pass validation"); @@ -287,6 +307,7 @@ mod tests { url: "rediss://:password@redis.example.com:6380".to_string(), pool_max: 8, retention_days: None, + schema: None, }; cfg.validate() .expect("rediss:// scheme should also be accepted"); @@ -298,6 +319,7 @@ mod tests { url: String::new(), pool_max: 16, retention_days: Some(30), + schema: None, }; let err = cfg.validate().expect_err("empty URL should fail"); assert!( @@ -312,6 +334,7 @@ mod tests { url: "http://localhost:6379".to_string(), pool_max: 16, retention_days: Some(30), + schema: None, }; let err = cfg.validate().expect_err("http scheme should be rejected"); assert!( @@ -326,6 +349,7 @@ mod tests { url: "redis:///0".to_string(), pool_max: 16, retention_days: Some(30), + schema: None, }; let err = cfg.validate().expect_err("missing host should fail"); assert!( @@ -340,6 +364,7 @@ mod tests { url: "redis://localhost:6379".to_string(), pool_max: 0, retention_days: Some(30), + schema: None, }; let err = cfg.validate().expect_err("pool_max=0 should fail"); assert!( diff --git a/data_connector/src/core.rs b/data_connector/src/core.rs index 55b5be803a..60bc571e9f 100644 --- a/data_connector/src/core.rs +++ b/data_connector/src/core.rs @@ -9,7 +9,7 @@ // 3. Response types + trait use std::{ - collections::HashMap, + collections::{HashMap, HashSet}, fmt::{Display, Formatter, Write}, }; @@ -465,13 +465,53 @@ pub trait ResponseStorage: Send + Sync { /// Delete a response async fn delete_response(&self, response_id: &ResponseId) -> ResponseResult<()>; - /// Get the chain of responses leading to a given response - /// Returns responses in chronological order (oldest first) + /// Get the chain of responses leading to a given response. + /// + /// Walks `previous_response_id` links from the given response backwards, + /// collecting up to `max_depth` responses (or unlimited if `None`). + /// Returns responses in chronological order (oldest first). + /// + /// The default implementation calls `self.get_response()` in a loop with + /// cycle detection to prevent infinite loops from self-referencing chains. + /// Backends that can walk the chain more efficiently (e.g. with a single + /// lock or a recursive SQL query) should override this. async fn get_response_chain( &self, response_id: &ResponseId, max_depth: Option, - ) -> ResponseResult; + ) -> ResponseResult { + let mut chain = ResponseChain::new(); + let mut current_id = Some(response_id.clone()); + let mut seen = HashSet::new(); + + while let Some(ref lookup_id) = current_id { + if let Some(limit) = max_depth { + if seen.len() >= limit { + break; + } + } + + // Cycle detection: error if we've already visited this ID. + if !seen.insert(lookup_id.clone()) { + return Err(ResponseStorageError::InvalidChain(format!( + "cycle detected at response {}", + lookup_id.0 + ))); + } + + let fetched = self.get_response(lookup_id).await?; + match fetched { + Some(response) => { + current_id.clone_from(&response.previous_response_id); + chain.responses.push(response); + } + None => break, + } + } + + chain.responses.reverse(); + Ok(chain) + } /// List recent responses for a safety identifier async fn list_identifier_responses( diff --git a/data_connector/src/lib.rs b/data_connector/src/lib.rs index 50c74c5cb9..c08269c538 100644 --- a/data_connector/src/lib.rs +++ b/data_connector/src/lib.rs @@ -21,6 +21,7 @@ mod noop; mod oracle; mod postgres; mod redis; +pub mod schema; // Re-export config types // Re-export core types and traits @@ -35,3 +36,5 @@ pub use config::{HistoryBackend, OracleConfig, PostgresConfig, RedisConfig}; pub use factory::{create_storage, StorageFactoryConfig}; // Re-export memory implementations for testing pub use memory::{MemoryConversationItemStorage, MemoryConversationStorage, MemoryResponseStorage}; +// Re-export schema config types +pub use schema::{SchemaConfig, TableConfig}; diff --git a/data_connector/src/oracle.rs b/data_connector/src/oracle.rs index f003fd2a29..a0efd18f43 100644 --- a/data_connector/src/oracle.rs +++ b/data_connector/src/oracle.rs @@ -21,19 +21,23 @@ use super::core::{ make_item_id, Conversation, ConversationId, ConversationItem, ConversationItemId, ConversationItemStorage, ConversationItemStorageError, ConversationMetadata, ConversationStorage, ConversationStorageError, ListParams, NewConversation, - NewConversationItem, ResponseChain, ResponseId, ResponseStorage, ResponseStorageError, - SortOrder, StoredResponse, + NewConversationItem, ResponseId, ResponseStorage, ResponseStorageError, SortOrder, + StoredResponse, }; use crate::{ - common::{parse_json_value, parse_metadata, parse_raw_response, parse_tool_calls}, + common::{ + build_response_select_base, parse_json_value, parse_metadata, parse_raw_response, + parse_tool_calls, + }, config::OracleConfig, + schema::SchemaConfig, }; // ============================================================================ // PART 1: OracleStore Helper + Common Utilities // ============================================================================ /// Schema initializer function signature for Oracle storage backends. -pub(crate) type SchemaInitFn = fn(&Connection) -> Result<(), String>; +pub(crate) type SchemaInitFn = fn(&Connection, &SchemaConfig) -> Result<(), String>; /// Shared Oracle connection pool infrastructure. /// @@ -41,6 +45,7 @@ pub(crate) type SchemaInitFn = fn(&Connection) -> Result<(), String>; /// It handles connection pooling, error mapping, and client configuration. pub(crate) struct OracleStore { pool: Pool, + pub(crate) schema: Arc, } impl OracleStore { @@ -49,6 +54,15 @@ impl OracleStore { /// Accepts a list of schema initializers that run on a single connection /// before the pool is created, ensuring all tables exist. pub fn new(config: &OracleConfig, init_schemas: &[SchemaInitFn]) -> Result { + // Extract and validate schema config. + // Oracle folds unquoted identifiers to uppercase, so existing tables + // and columns are CONVERSATIONS, CONV_ID, etc. Uppercase all + // configured names so that quoted references match reality. + let mut schema = config.schema.clone().unwrap_or_default(); + schema.uppercase_for_oracle(); + schema.validate()?; + let schema = Arc::new(schema); + // Configure Oracle client (wallet env vars, etc.) configure_oracle_env(config)?; @@ -62,7 +76,7 @@ impl OracleStore { .map_err(map_oracle_error)?; for init_schema in init_schemas { - init_schema(&conn)?; + init_schema(&conn, &schema)?; } drop(conn); @@ -84,7 +98,7 @@ impl OracleStore { .build() .map_err(|e| format!("Failed to build Oracle pool: {e}"))?; - Ok(Self { pool }) + Ok(Self { pool, schema }) } /// Execute function with a connection from the pool @@ -113,6 +127,7 @@ impl Clone for OracleStore { fn clone(&self) -> Self { Self { pool: self.pool.clone(), + schema: self.schema.clone(), } } } @@ -266,21 +281,30 @@ impl OracleConversationStorage { Self { store } } - pub(crate) fn init_schema(conn: &Connection) -> Result<(), String> { + pub(crate) fn init_schema(conn: &Connection, schema: &SchemaConfig) -> Result<(), String> { + let s = &schema.conversations; + // Table and column names are already uppercased by OracleStore::new(). + let table = s.qualified_table(schema.owner.as_deref()); + let exists: i64 = conn .query_row_as( - "SELECT COUNT(*) FROM user_tables WHERE table_name = 'CONVERSATIONS'", + &format!( + "SELECT COUNT(*) FROM user_tables WHERE table_name = '{}'", + s.table + ), &[], ) .map_err(map_oracle_error)?; if exists == 0 { + let col_defs = [ + format!("{} VARCHAR2(64) PRIMARY KEY", s.col("id")), + format!("{} TIMESTAMP WITH TIME ZONE", s.col("created_at")), + format!("{} CLOB", s.col("metadata")), + ]; + conn.execute( - "CREATE TABLE conversations ( - id VARCHAR2(64) PRIMARY KEY, - created_at TIMESTAMP WITH TIME ZONE, - metadata CLOB - )", + &format!("CREATE TABLE {table} ({})", col_defs.join(", ")), &[], ) .map_err(map_oracle_error)?; @@ -289,14 +313,6 @@ impl OracleConversationStorage { Ok(()) } - /// Parse raw metadata JSON into `ConversationMetadata`. - /// - /// Delegates to the shared `parse_conversation_metadata` in `common.rs`. - /// The previous Oracle-specific version performed an extra `Value` type-check - /// (rejecting non-object JSON explicitly), but the common function achieves the - /// same result: `serde_json::from_str::` already rejects non-object JSON - /// with a descriptive error, and metadata is always serialized from a `JsonMap` by - /// our own code, so non-object values cannot occur in practice. fn parse_metadata( raw: Option, ) -> Result, ConversationStorageError> { @@ -319,15 +335,29 @@ impl ConversationStorage for OracleConversationStorage { .as_ref() .map(serde_json::to_string) .transpose()?; + let schema = self.store.schema.clone(); self.store .execute(move |conn| { - conn.execute( - "INSERT INTO conversations (id, created_at, metadata) VALUES (:1, :2, :3)", - &[&id_str, &created_at, &metadata_json], - ) - .map(|_| ()) - .map_err(map_oracle_error) + let s = &schema.conversations; + let table = s.qualified_table(schema.owner.as_deref()); + let col_id = s.col("id"); + let col_created = s.col("created_at"); + let col_meta = s.col("metadata"); + + let columns = [col_id, col_created, col_meta]; + let placeholders: Vec = + (1..=columns.len()).map(|i| format!(":{i}")).collect(); + let params: Vec<&dyn ToSql> = vec![&id_str, &created_at, &metadata_json]; + + let sql = format!( + "INSERT INTO {table} ({}) VALUES ({})", + columns.join(", "), + placeholders.join(", ") + ); + conn.execute(&sql, ¶ms[..]) + .map(|_| ()) + .map_err(map_oracle_error) }) .await .map_err(ConversationStorageError::StorageError)?; @@ -340,19 +370,29 @@ impl ConversationStorage for OracleConversationStorage { id: &ConversationId, ) -> Result, ConversationStorageError> { let lookup = id.0.clone(); + let schema = self.store.schema.clone(); + self.store .execute(move |conn| { - let mut stmt = conn - .statement("SELECT id, created_at, metadata FROM conversations WHERE id = :1") - .build() - .map_err(map_oracle_error)?; + let s = &schema.conversations; + let table = s.qualified_table(schema.owner.as_deref()); + let col_id = s.col("id"); + let col_created = s.col("created_at"); + let col_meta = s.col("metadata"); + + let sql = format!( + "SELECT {col_id}, {col_created}, {col_meta} FROM {table} WHERE {col_id} = :1" + ); + let mut stmt = conn.statement(&sql).build().map_err(map_oracle_error)?; let mut rows = stmt.query(&[&lookup]).map_err(map_oracle_error)?; if let Some(row_res) = rows.next() { let row = row_res.map_err(map_oracle_error)?; - let id: String = row.get(0).map_err(map_oracle_error)?; - let created_at: DateTime = row.get(1).map_err(map_oracle_error)?; - let metadata_raw: Option = row.get(2).map_err(map_oracle_error)?; + let id: String = row.get(col_id).map_err(map_oracle_error)?; + let created_at: DateTime = + row.get(col_created).map_err(map_oracle_error)?; + let metadata_raw: Option = + row.get(col_meta).map_err(map_oracle_error)?; let metadata = Self::parse_metadata(metadata_raw).map_err(|e| e.to_string())?; Ok(Some(Conversation::with_parts( ConversationId(id), @@ -375,18 +415,22 @@ impl ConversationStorage for OracleConversationStorage { let id_str = id.0.clone(); let metadata_json = metadata.as_ref().map(serde_json::to_string).transpose()?; let conversation_id = id.clone(); + let schema = self.store.schema.clone(); self.store .execute(move |conn| { - let mut stmt = conn - .statement( - "UPDATE conversations \ - SET metadata = :1 \ - WHERE id = :2 \ - RETURNING created_at INTO :3", - ) - .build() - .map_err(map_oracle_error)?; + let s = &schema.conversations; + let table = s.qualified_table(schema.owner.as_deref()); + let col_id = s.col("id"); + let col_meta = s.col("metadata"); + let col_created = s.col("created_at"); + + let sql = format!( + "UPDATE {table} SET {col_meta} = :1 \ + WHERE {col_id} = :2 \ + RETURNING {col_created} INTO :3" + ); + let mut stmt = conn.statement(&sql).build().map_err(map_oracle_error)?; stmt.bind(3, &OracleType::TimestampTZ(6)) .map_err(map_oracle_error)?; @@ -418,11 +462,20 @@ impl ConversationStorage for OracleConversationStorage { id: &ConversationId, ) -> Result { let id_str = id.0.clone(); + let schema = self.store.schema.clone(); + let res = self .store .execute(move |conn| { - conn.execute("DELETE FROM conversations WHERE id = :1", &[&id_str]) - .map_err(map_oracle_error) + let s = &schema.conversations; + let table = s.qualified_table(schema.owner.as_deref()); + let col_id = s.col("id"); + + conn.execute( + &format!("DELETE FROM {table} WHERE {col_id} = :1"), + &[&id_str], + ) + .map_err(map_oracle_error) }) .await .map_err(ConversationStorageError::StorageError)?; @@ -448,51 +501,75 @@ impl OracleConversationItemStorage { Self { store } } - pub(crate) fn init_schema(conn: &Connection) -> Result<(), String> { + pub(crate) fn init_schema(conn: &Connection, schema: &SchemaConfig) -> Result<(), String> { + let si = &schema.conversation_items; + // Table and column names are already uppercased by OracleStore::new(). + let si_table = si.qualified_table(schema.owner.as_deref()); + let exists_items: i64 = conn .query_row_as( - "SELECT COUNT(*) FROM user_tables WHERE table_name = 'CONVERSATION_ITEMS'", + &format!( + "SELECT COUNT(*) FROM user_tables WHERE table_name = '{}'", + si.table + ), &[], ) .map_err(map_oracle_error)?; if exists_items == 0 { + let col_defs = [ + format!("{} VARCHAR2(64) PRIMARY KEY", si.col("id")), + format!("{} VARCHAR2(64)", si.col("response_id")), + format!("{} VARCHAR2(32) NOT NULL", si.col("item_type")), + format!("{} VARCHAR2(32)", si.col("role")), + format!("{} CLOB", si.col("content")), + format!("{} VARCHAR2(32)", si.col("status")), + format!("{} TIMESTAMP WITH TIME ZONE", si.col("created_at")), + ]; + conn.execute( - "CREATE TABLE conversation_items ( - id VARCHAR2(64) PRIMARY KEY, - response_id VARCHAR2(64), - item_type VARCHAR2(32) NOT NULL, - role VARCHAR2(32), - content CLOB, - status VARCHAR2(32), - created_at TIMESTAMP WITH TIME ZONE - )", + &format!("CREATE TABLE {si_table} ({})", col_defs.join(", ")), &[], ) .map_err(map_oracle_error)?; } + let sl = &schema.conversation_item_links; + let sl_table = sl.qualified_table(schema.owner.as_deref()); + let exists_links: i64 = conn .query_row_as( - "SELECT COUNT(*) FROM user_tables WHERE table_name = 'CONVERSATION_ITEM_LINKS'", + &format!( + "SELECT COUNT(*) FROM user_tables WHERE table_name = '{}'", + sl.table + ), &[], ) .map_err(map_oracle_error)?; if exists_links == 0 { + let col_cid = sl.col("conversation_id"); + let col_iid = sl.col("item_id"); + let col_added = sl.col("added_at"); + + let pk_name = format!("PK_{}", sl.table); + let idx_name = format!("{}_CONV_IDX", sl.table); + + let col_defs = [ + format!("{col_cid} VARCHAR2(64) NOT NULL"), + format!("{col_iid} VARCHAR2(64) NOT NULL"), + format!("{col_added} TIMESTAMP WITH TIME ZONE"), + format!("CONSTRAINT {pk_name} PRIMARY KEY ({col_cid}, {col_iid})"), + ]; + conn.execute( - "CREATE TABLE conversation_item_links ( - conversation_id VARCHAR2(64) NOT NULL, - item_id VARCHAR2(64) NOT NULL, - added_at TIMESTAMP WITH TIME ZONE, - CONSTRAINT pk_conv_item_link PRIMARY KEY (conversation_id, item_id) - )", + &format!("CREATE TABLE {sl_table} ({})", col_defs.join(", ")), &[], ) .map_err(map_oracle_error)?; conn.execute( - "CREATE INDEX conv_item_links_conv_idx ON conversation_item_links (conversation_id, added_at)", + &format!("CREATE INDEX {idx_name} ON {sl_table} ({col_cid}, {col_added})"), &[], ) .map_err(map_oracle_error)?; @@ -520,22 +597,52 @@ impl ConversationItemStorage for OracleConversationItemStorage { let created_at = Utc::now(); let content_json = serde_json::to_string(&content)?; - // Clone fields needed for both the closure and the return value. - // The closure consumes the clones; originals go into the return struct. let id_str = id.0.clone(); let cl_response_id = response_id.clone(); let cl_item_type = item_type.clone(); let cl_role = role.clone(); let cl_status = status.clone(); + let schema = self.store.schema.clone(); self.store .execute(move |conn| { - conn.execute( - "INSERT INTO conversation_items (id, response_id, item_type, role, content, status, created_at) \ - VALUES (:1, :2, :3, :4, :5, :6, :7)", - &[&id_str, &cl_response_id, &cl_item_type, &cl_role, &content_json, &cl_status, &created_at], - ) - .map_err(map_oracle_error)?; + let si = &schema.conversation_items; + let table = si.qualified_table(schema.owner.as_deref()); + let col_id = si.col("id"); + let col_resp = si.col("response_id"); + let col_type = si.col("item_type"); + let col_role = si.col("role"); + let col_content = si.col("content"); + let col_status = si.col("status"); + let col_created = si.col("created_at"); + + let columns = [ + col_id, + col_resp, + col_type, + col_role, + col_content, + col_status, + col_created, + ]; + let placeholders: Vec = + (1..=columns.len()).map(|i| format!(":{i}")).collect(); + let params: Vec<&dyn ToSql> = vec![ + &id_str, + &cl_response_id, + &cl_item_type, + &cl_role, + &content_json, + &cl_status, + &created_at, + ]; + + let sql = format!( + "INSERT INTO {table} ({}) VALUES ({})", + columns.join(", "), + placeholders.join(", ") + ); + conn.execute(&sql, ¶ms[..]).map_err(map_oracle_error)?; Ok(()) }) .await @@ -560,13 +667,21 @@ impl ConversationItemStorage for OracleConversationItemStorage { ) -> Result<(), ConversationItemStorageError> { let cid = conversation_id.0.clone(); let iid = item_id.0.clone(); + let schema = self.store.schema.clone(); + self.store .execute(move |conn| { - conn.execute( - "INSERT INTO conversation_item_links (conversation_id, item_id, added_at) VALUES (:1, :2, :3)", - &[&cid, &iid, &added_at], - ) - .map_err(map_oracle_error)?; + let sl = &schema.conversation_item_links; + let table = sl.qualified_table(schema.owner.as_deref()); + let col_cid = sl.col("conversation_id"); + let col_iid = sl.col("item_id"); + let col_added = sl.col("added_at"); + + let sql = format!( + "INSERT INTO {table} ({col_cid}, {col_iid}, {col_added}) VALUES (:1, :2, :3)" + ); + conn.execute(&sql, &[&cid, &iid, &added_at]) + .map_err(map_oracle_error)?; Ok(()) }) .await @@ -582,24 +697,31 @@ impl ConversationItemStorage for OracleConversationItemStorage { let limit: i64 = params.limit as i64; let order_desc = matches!(params.order, SortOrder::Desc); let after_id = params.after.clone(); + let schema = self.store.schema.clone(); // Resolve the added_at of the after cursor if provided let after_key: Option<(DateTime, String)> = if let Some(ref aid) = after_id { + let schema2 = schema.clone(); self.store .execute({ let cid = cid.clone(); let aid = aid.clone(); move |conn| { - let mut stmt = conn - .statement( - "SELECT added_at FROM conversation_item_links WHERE conversation_id = :1 AND item_id = :2", - ) - .build() - .map_err(map_oracle_error)?; + let sl = &schema2.conversation_item_links; + let table = sl.qualified_table(schema2.owner.as_deref()); + let col_added = sl.col("added_at"); + let col_cid = sl.col("conversation_id"); + let col_iid = sl.col("item_id"); + + let sql = format!( + "SELECT {col_added} FROM {table} \ + WHERE {col_cid} = :1 AND {col_iid} = :2" + ); + let mut stmt = conn.statement(&sql).build().map_err(map_oracle_error)?; let mut rows = stmt.query(&[&cid, &aid]).map_err(map_oracle_error)?; if let Some(row_res) = rows.next() { let row = row_res.map_err(map_oracle_error)?; - let ts: DateTime = row.get(0).map_err(map_oracle_error)?; + let ts: DateTime = row.get(col_added).map_err(map_oracle_error)?; Ok(Some((ts, aid))) } else { Ok(None) @@ -612,86 +734,97 @@ impl ConversationItemStorage for OracleConversationItemStorage { None }; - // Build the main list query - let rows: Vec<(String, Option, String, Option, Option, Option, DateTime)> = - self.store - .execute({ - let cid = cid.clone(); - move |conn| { - let mut sql = String::from( - "SELECT i.id, i.response_id, i.item_type, i.role, i.content, i.status, i.created_at \ - FROM conversation_item_links l \ - JOIN conversation_items i ON i.id = l.item_id \ - WHERE l.conversation_id = :cid", - ); - - // Cursor predicate - if let Some((_ts, _iid)) = &after_key { - if order_desc { - sql.push_str(" AND (l.added_at < :ats OR (l.added_at = :ats AND l.item_id < :iid))"); - } else { - sql.push_str(" AND (l.added_at > :ats OR (l.added_at = :ats AND l.item_id > :iid))"); - } - } - - // Order and limit + // Build the main list query and construct items directly in the closure. + self.store + .execute({ + let cid = cid.clone(); + let schema = schema.clone(); + move |conn| { + let si = &schema.conversation_items; + let sl = &schema.conversation_item_links; + let si_table = si.qualified_table(schema.owner.as_deref()); + let sl_table = sl.qualified_table(schema.owner.as_deref()); + let si_col_id = si.col("id"); + let si_col_resp = si.col("response_id"); + let si_col_type = si.col("item_type"); + let si_col_role = si.col("role"); + let si_col_content = si.col("content"); + let si_col_status = si.col("status"); + let si_col_created = si.col("created_at"); + let sl_col_cid = sl.col("conversation_id"); + let sl_col_iid = sl.col("item_id"); + let sl_col_added = sl.col("added_at"); + + let mut sql = format!( + "SELECT i.{si_col_id}, i.{si_col_resp}, i.{si_col_type}, \ + i.{si_col_role}, i.{si_col_content}, i.{si_col_status}, i.{si_col_created} \ + FROM {sl_table} l \ + JOIN {si_table} i ON i.{si_col_id} = l.{sl_col_iid} \ + WHERE l.{sl_col_cid} = :cid" + ); + + if let Some((_ts, _iid)) = &after_key { if order_desc { - sql.push_str(" ORDER BY l.added_at DESC, l.item_id DESC"); + sql.push_str(&format!( + " AND (l.{sl_col_added} < :ats OR \ + (l.{sl_col_added} = :ats AND l.{sl_col_iid} < :iid))" + )); } else { - sql.push_str(" ORDER BY l.added_at ASC, l.item_id ASC"); + sql.push_str(&format!( + " AND (l.{sl_col_added} > :ats OR \ + (l.{sl_col_added} = :ats AND l.{sl_col_iid} > :iid))" + )); } - sql.push_str(" FETCH NEXT :limit ROWS ONLY"); - - // Build params and perform a named SELECT query - let mut params_vec: Vec<(&str, &dyn ToSql)> = vec![("cid", &cid)]; - if let Some((ts, iid)) = &after_key { - params_vec.push(("ats", ts)); - params_vec.push(("iid", iid)); - } - params_vec.push(("limit", &limit)); - - let rows_iter = conn.query_named(&sql, ¶ms_vec).map_err(map_oracle_error)?; + } - let mut out = Vec::new(); - for row_res in rows_iter { - let row = row_res.map_err(map_oracle_error)?; - let id: String = row.get(0).map_err(map_oracle_error)?; - let resp_id: Option = row.get(1).map_err(map_oracle_error)?; - let item_type: String = row.get(2).map_err(map_oracle_error)?; - let role: Option = row.get(3).map_err(map_oracle_error)?; - let content_raw: Option = row.get(4).map_err(map_oracle_error)?; - let status: Option = row.get(5).map_err(map_oracle_error)?; - let created_at: DateTime = row.get(6).map_err(map_oracle_error)?; - out.push((id, resp_id, item_type, role, content_raw, status, created_at)); - } - Ok(out) + if order_desc { + sql.push_str(&format!( + " ORDER BY l.{sl_col_added} DESC, l.{sl_col_iid} DESC" + )); + } else { + sql.push_str(&format!( + " ORDER BY l.{sl_col_added} ASC, l.{sl_col_iid} ASC" + )); } - }) - .await - .map_err(ConversationItemStorageError::StorageError)?; + sql.push_str(" FETCH NEXT :limit ROWS ONLY"); - // Map rows to ConversationItem - rows.into_iter() - .map( - |(id, resp_id, item_type, role, content_raw, status, created_at)| { - let content = match content_raw { - Some(s) => { - serde_json::from_str(&s).map_err(ConversationItemStorageError::from)? - } - None => Value::Null, - }; - Ok(ConversationItem { - id: ConversationItemId(id), - response_id: resp_id, - item_type, - role, - content, - status, - created_at, - }) - }, - ) - .collect() + let mut params_vec: Vec<(&str, &dyn ToSql)> = vec![("cid", &cid)]; + if let Some((ts, iid)) = &after_key { + params_vec.push(("ats", ts)); + params_vec.push(("iid", iid)); + } + params_vec.push(("limit", &limit)); + + let rows_iter = + conn.query_named(&sql, ¶ms_vec).map_err(map_oracle_error)?; + + let mut items = Vec::new(); + for row_res in rows_iter { + let row = row_res.map_err(map_oracle_error)?; + let content_raw: Option = + row.get(si_col_content).map_err(map_oracle_error)?; + let content: Value = match content_raw { + Some(s) => serde_json::from_str(&s).map_err(|e| e.to_string())?, + None => Value::Null, + }; + + items.push(ConversationItem { + id: ConversationItemId( + row.get(si_col_id).map_err(map_oracle_error)?, + ), + response_id: row.get(si_col_resp).map_err(map_oracle_error)?, + item_type: row.get(si_col_type).map_err(map_oracle_error)?, + role: row.get(si_col_role).map_err(map_oracle_error)?, + content, + status: row.get(si_col_status).map_err(map_oracle_error)?, + created_at: row.get(si_col_created).map_err(map_oracle_error)?, + }); + } + Ok(items) + } + }) + .await + .map_err(ConversationItemStorageError::StorageError) } async fn get_item( @@ -699,28 +832,40 @@ impl ConversationItemStorage for OracleConversationItemStorage { item_id: &ConversationItemId, ) -> Result, ConversationItemStorageError> { let iid = item_id.0.clone(); + let schema = self.store.schema.clone(); self.store .execute(move |conn| { - let mut stmt = conn - .statement( - "SELECT id, response_id, item_type, role, content, status, created_at \ - FROM conversation_items WHERE id = :1", - ) - .build() - .map_err(map_oracle_error)?; - + let si = &schema.conversation_items; + let table = si.qualified_table(schema.owner.as_deref()); + let col_id = si.col("id"); + let col_resp = si.col("response_id"); + let col_type = si.col("item_type"); + let col_role = si.col("role"); + let col_content = si.col("content"); + let col_status = si.col("status"); + let col_created = si.col("created_at"); + + let sql = format!( + "SELECT {col_id}, {col_resp}, {col_type}, {col_role}, \ + {col_content}, {col_status}, {col_created} \ + FROM {table} WHERE {col_id} = :1" + ); + let mut stmt = conn.statement(&sql).build().map_err(map_oracle_error)?; let mut rows = stmt.query(&[&iid]).map_err(map_oracle_error)?; if let Some(row_res) = rows.next() { let row = row_res.map_err(map_oracle_error)?; - let id: String = row.get(0).map_err(map_oracle_error)?; - let response_id: Option = row.get(1).map_err(map_oracle_error)?; - let item_type: String = row.get(2).map_err(map_oracle_error)?; - let role: Option = row.get(3).map_err(map_oracle_error)?; - let content_raw: Option = row.get(4).map_err(map_oracle_error)?; - let status: Option = row.get(5).map_err(map_oracle_error)?; - let created_at: DateTime = row.get(6).map_err(map_oracle_error)?; + let id: String = row.get(col_id).map_err(map_oracle_error)?; + let response_id: Option = + row.get(col_resp).map_err(map_oracle_error)?; + let item_type: String = row.get(col_type).map_err(map_oracle_error)?; + let role: Option = row.get(col_role).map_err(map_oracle_error)?; + let content_raw: Option = + row.get(col_content).map_err(map_oracle_error)?; + let status: Option = row.get(col_status).map_err(map_oracle_error)?; + let created_at: DateTime = + row.get(col_created).map_err(map_oracle_error)?; let content = match content_raw { Some(s) => serde_json::from_str(&s).map_err(|e| e.to_string())?, @@ -751,14 +896,19 @@ impl ConversationItemStorage for OracleConversationItemStorage { ) -> Result { let cid = conversation_id.0.clone(); let iid = item_id.0.clone(); + let schema = self.store.schema.clone(); self.store .execute(move |conn| { + let sl = &schema.conversation_item_links; + let table = sl.qualified_table(schema.owner.as_deref()); + let col_cid = sl.col("conversation_id"); + let col_iid = sl.col("item_id"); + + let sql = + format!("SELECT COUNT(*) FROM {table} WHERE {col_cid} = :1 AND {col_iid} = :2"); let count: i64 = conn - .query_row_as( - "SELECT COUNT(*) FROM conversation_item_links WHERE conversation_id = :1 AND item_id = :2", - &[&cid, &iid], - ) + .query_row_as(&sql, &[&cid, &iid]) .map_err(map_oracle_error)?; Ok(count > 0) }) @@ -773,11 +923,17 @@ impl ConversationItemStorage for OracleConversationItemStorage { ) -> Result<(), ConversationItemStorageError> { let cid = conversation_id.0.clone(); let iid = item_id.0.clone(); + let schema = self.store.schema.clone(); self.store .execute(move |conn| { + let sl = &schema.conversation_item_links; + let table = sl.qualified_table(schema.owner.as_deref()); + let col_cid = sl.col("conversation_id"); + let col_iid = sl.col("item_id"); + conn.execute( - "DELETE FROM conversation_item_links WHERE conversation_id = :1 AND item_id = :2", + &format!("DELETE FROM {table} WHERE {col_cid} = :1 AND {col_iid} = :2"), &[&cid, &iid], ) .map_err(map_oracle_error)?; @@ -792,82 +948,113 @@ impl ConversationItemStorage for OracleConversationItemStorage { // PART 4: OracleResponseStorage // ============================================================================ -const SELECT_BASE: &str = "SELECT id, previous_response_id, input, instructions, output, \ - tool_calls, metadata, created_at, safety_identifier, model, conversation_id, raw_response FROM responses"; - #[derive(Clone)] pub(super) struct OracleResponseStorage { store: OracleStore, + select_base: String, } impl OracleResponseStorage { pub fn new(store: OracleStore) -> Self { - Self { store } + let select_base = build_response_select_base(&store.schema); + Self { store, select_base } } - pub(crate) fn init_schema(conn: &Connection) -> Result<(), String> { + pub(crate) fn init_schema(conn: &Connection, schema: &SchemaConfig) -> Result<(), String> { + let s = &schema.responses; + // Table and column names are already uppercased by OracleStore::new(). + let table = s.qualified_table(schema.owner.as_deref()); + let exists: i64 = conn .query_row_as( - "SELECT COUNT(*) FROM user_tables WHERE table_name = 'RESPONSES'", + &format!( + "SELECT COUNT(*) FROM user_tables WHERE table_name = '{}'", + s.table + ), &[], ) .map_err(map_oracle_error)?; if exists == 0 { + let col_defs = vec![ + format!("{} VARCHAR2(64) PRIMARY KEY", s.col("id")), + format!("{} VARCHAR2(64)", s.col("conversation_id")), + format!("{} VARCHAR2(64)", s.col("previous_response_id")), + format!("{} CLOB", s.col("input")), + format!("{} CLOB", s.col("instructions")), + format!("{} CLOB", s.col("output")), + format!("{} CLOB", s.col("tool_calls")), + format!("{} CLOB", s.col("metadata")), + format!("{} TIMESTAMP WITH TIME ZONE", s.col("created_at")), + format!("{} VARCHAR2(128)", s.col("safety_identifier")), + format!("{} VARCHAR2(128)", s.col("model")), + format!("{} CLOB", s.col("raw_response")), + ]; + conn.execute( - "CREATE TABLE responses ( - id VARCHAR2(64) PRIMARY KEY, - conversation_id VARCHAR2(64), - previous_response_id VARCHAR2(64), - input CLOB, - instructions CLOB, - output CLOB, - tool_calls CLOB, - metadata CLOB, - created_at TIMESTAMP WITH TIME ZONE, - safety_identifier VARCHAR2(128), - model VARCHAR2(128), - raw_response CLOB - )", + &format!("CREATE TABLE {table} ({})", col_defs.join(", ")), &[], ) .map_err(map_oracle_error)?; } else { - Self::alter_safety_identifier_column(conn)?; - Self::remove_user_id_column_if_exists(conn)?; + Self::alter_safety_identifier_column(conn, schema)?; + Self::remove_user_id_column_if_exists(conn, schema)?; } + let prev = s.col("previous_response_id"); + let prev_idx = format!("{}_PREV_IDX", s.table); create_index_if_missing( conn, - "RESPONSES_PREV_IDX", - "CREATE INDEX responses_prev_idx ON responses(previous_response_id)", + &s.table, + &prev_idx, + &format!("CREATE INDEX {prev_idx} ON {table}({prev})"), )?; + + let safety = s.col("safety_identifier"); + let user_idx = format!("{}_USER_IDX", s.table); create_index_if_missing( conn, - "RESPONSES_USER_IDX", - "CREATE INDEX responses_user_idx ON responses(safety_identifier)", + &s.table, + &user_idx, + &format!("CREATE INDEX {user_idx} ON {table}({safety})"), )?; Ok(()) } - // Alter safety_identifier column if missing - fn alter_safety_identifier_column(conn: &Connection) -> Result<(), String> { + fn alter_safety_identifier_column( + conn: &Connection, + schema: &SchemaConfig, + ) -> Result<(), String> { + let s = &schema.responses; + let col_safety = s.col("safety_identifier"); + // Table and column names are already uppercased by OracleStore::new(). + let col_upper = col_safety.to_uppercase(); + let table = s.qualified_table(schema.owner.as_deref()); + let present: i64 = conn .query_row_as( - "SELECT COUNT(*) FROM user_tab_columns WHERE table_name = 'RESPONSES' AND column_name = 'SAFETY_IDENTIFIER'", + &format!( + "SELECT COUNT(*) FROM user_tab_columns \ + WHERE table_name = '{}' AND column_name = '{col_upper}'", + s.table + ), &[], ) .map_err(map_oracle_error)?; if present == 0 { if let Err(err) = conn.execute( - "ALTER TABLE responses ADD (safety_identifier VARCHAR2(128))", + &format!("ALTER TABLE {table} ADD ({col_safety} VARCHAR2(128))"), &[], ) { let present_after: i64 = conn .query_row_as( - "SELECT COUNT(*) FROM user_tab_columns WHERE table_name = 'RESPONSES' AND column_name = 'SAFETY_IDENTIFIER'", + &format!( + "SELECT COUNT(*) FROM user_tab_columns \ + WHERE table_name = '{}' AND column_name = '{col_upper}'", + s.table + ), &[], ) .map_err(map_oracle_error)?; @@ -880,20 +1067,35 @@ impl OracleResponseStorage { Ok(()) } - // Remove user_id column if exists - fn remove_user_id_column_if_exists(conn: &Connection) -> Result<(), String> { + fn remove_user_id_column_if_exists( + conn: &Connection, + schema: &SchemaConfig, + ) -> Result<(), String> { + // Table and column names are already uppercased by OracleStore::new(). + let s = &schema.responses; + let table = s.qualified_table(schema.owner.as_deref()); + let present: i64 = conn .query_row_as( - "SELECT COUNT(*) FROM user_tab_columns WHERE table_name = 'RESPONSES' AND column_name = 'USER_ID'", + &format!( + "SELECT COUNT(*) FROM user_tab_columns \ + WHERE table_name = '{}' AND column_name = 'USER_ID'", + s.table + ), &[], ) .map_err(map_oracle_error)?; if present > 0 { - if let Err(err) = conn.execute("ALTER TABLE responses DROP COLUMN USER_ID", &[]) { + if let Err(err) = conn.execute(&format!("ALTER TABLE {table} DROP COLUMN USER_ID"), &[]) + { let present_after: i64 = conn .query_row_as( - "SELECT COUNT(*) FROM user_tab_columns WHERE table_name = 'RESPONSES' AND column_name = 'USER_ID'", + &format!( + "SELECT COUNT(*) FROM user_tab_columns \ + WHERE table_name = '{}' AND column_name = 'USER_ID'", + s.table + ), &[], ) .map_err(map_oracle_error)?; @@ -906,19 +1108,33 @@ impl OracleResponseStorage { Ok(()) } - fn build_response_from_row(row: &Row) -> Result { - let id: String = row.get(0).map_err(map_oracle_error)?; - let previous: Option = row.get(1).map_err(map_oracle_error)?; - let input_json: Option = row.get(2).map_err(map_oracle_error)?; - let instructions: Option = row.get(3).map_err(map_oracle_error)?; - let output_json: Option = row.get(4).map_err(map_oracle_error)?; - let tool_calls_json: Option = row.get(5).map_err(map_oracle_error)?; - let metadata_json: Option = row.get(6).map_err(map_oracle_error)?; - let created_at: DateTime = row.get(7).map_err(map_oracle_error)?; - let safety_identifier: Option = row.get(8).map_err(map_oracle_error)?; - let model: Option = row.get(9).map_err(map_oracle_error)?; - let conversation_id: Option = row.get(10).map_err(map_oracle_error)?; - let raw_response_json: Option = row.get(11).map_err(map_oracle_error)?; + fn build_response_from_row(row: &Row, schema: &SchemaConfig) -> Result { + let s = &schema.responses; + let col_id = s.col("id"); + let col_created = s.col("created_at"); + + let id: String = row.get(col_id).map_err(map_oracle_error)?; + let created_at: DateTime = row.get(col_created).map_err(map_oracle_error)?; + + let previous: Option = row + .get(s.col("previous_response_id")) + .map_err(map_oracle_error)?; + let input_json: Option = row.get(s.col("input")).map_err(map_oracle_error)?; + let instructions: Option = + row.get(s.col("instructions")).map_err(map_oracle_error)?; + let output_json: Option = row.get(s.col("output")).map_err(map_oracle_error)?; + let tool_calls_json: Option = + row.get(s.col("tool_calls")).map_err(map_oracle_error)?; + let metadata_json: Option = row.get(s.col("metadata")).map_err(map_oracle_error)?; + let safety_identifier: Option = row + .get(s.col("safety_identifier")) + .map_err(map_oracle_error)?; + let model: Option = row.get(s.col("model")).map_err(map_oracle_error)?; + let conversation_id: Option = row + .get(s.col("conversation_id")) + .map_err(map_oracle_error)?; + let raw_response_json: Option = + row.get(s.col("raw_response")).map_err(map_oracle_error)?; let previous_response_id = previous.map(ResponseId); let tool_calls = parse_tool_calls(tool_calls_json)?; @@ -965,7 +1181,6 @@ impl ResponseStorage for OracleResponseStorage { raw_response, } = response; - // Clone only the return value; everything else moves into the closure. let return_id = id.clone(); let response_id_str = id.0; @@ -975,30 +1190,46 @@ impl ResponseStorage for OracleResponseStorage { let json_tool_calls = serde_json::to_string(&tool_calls)?; let json_metadata = serde_json::to_string(&metadata)?; let json_raw_response = serde_json::to_string(&raw_response)?; + let schema = self.store.schema.clone(); self.store .execute(move |conn| { - conn.execute( - "INSERT INTO responses (id, previous_response_id, input, instructions, output, \ - tool_calls, metadata, created_at, safety_identifier, model, conversation_id, raw_response) \ - VALUES (:1, :2, :3, :4, :5, :6, :7, :8, :9, :10, :11, :12)", - &[ - &response_id_str, - &previous_id, - &json_input, - &instructions, - &json_output, - &json_tool_calls, - &json_metadata, - &created_at, - &safety_identifier, - &model, - &conversation_id, - &json_raw_response, - ], - ) - .map(|_| ()) - .map_err(map_oracle_error) + let s = &schema.responses; + let table = s.qualified_table(schema.owner.as_deref()); + + // Build column list and placeholders dynamically + let logical_fields: &[(&str, &dyn ToSql)] = &[ + ("id", &response_id_str), + ("previous_response_id", &previous_id), + ("input", &json_input), + ("instructions", &instructions), + ("output", &json_output), + ("tool_calls", &json_tool_calls), + ("metadata", &json_metadata), + ("created_at", &created_at), + ("safety_identifier", &safety_identifier), + ("model", &model), + ("conversation_id", &conversation_id), + ("raw_response", &json_raw_response), + ]; + + let mut columns = Vec::new(); + let mut params: Vec<&dyn ToSql> = Vec::new(); + for &(logical, val) in logical_fields { + columns.push(s.col(logical)); + params.push(val); + } + + let placeholders: Vec = + (1..=params.len()).map(|i| format!(":{i}")).collect(); + let sql = format!( + "INSERT INTO {table} ({}) VALUES ({})", + columns.join(", "), + placeholders.join(", ") + ); + conn.execute(&sql, ¶ms[..]) + .map(|_| ()) + .map_err(map_oracle_error) }) .await .map_err(ResponseStorageError::StorageError)?; @@ -1011,17 +1242,19 @@ impl ResponseStorage for OracleResponseStorage { response_id: &ResponseId, ) -> Result, ResponseStorageError> { let id = response_id.0.clone(); + let select_base = self.select_base.clone(); + let schema = self.store.schema.clone(); + self.store .execute(move |conn| { - let mut stmt = conn - .statement(&format!("{SELECT_BASE} WHERE id = :1")) - .build() - .map_err(map_oracle_error)?; + let col_id = schema.responses.col("id"); + let sql = format!("{select_base} WHERE {col_id} = :1"); + let mut stmt = conn.statement(&sql).build().map_err(map_oracle_error)?; let mut rows = stmt.query(&[&id]).map_err(map_oracle_error)?; match rows.next() { Some(row) => { let row = row.map_err(map_oracle_error)?; - Self::build_response_from_row(&row).map(Some) + Self::build_response_from_row(&row, &schema).map(Some) } None => Ok(None), } @@ -1032,9 +1265,15 @@ impl ResponseStorage for OracleResponseStorage { async fn delete_response(&self, response_id: &ResponseId) -> Result<(), ResponseStorageError> { let id = response_id.0.clone(); + let schema = self.store.schema.clone(); + self.store .execute(move |conn| { - conn.execute("DELETE FROM responses WHERE id = :1", &[&id]) + let s = &schema.responses; + let table = s.qualified_table(schema.owner.as_deref()); + let col_id = s.col("id"); + + conn.execute(&format!("DELETE FROM {table} WHERE {col_id} = :1"), &[&id]) .map(|_| ()) .map_err(map_oracle_error) }) @@ -1042,52 +1281,28 @@ impl ResponseStorage for OracleResponseStorage { .map_err(ResponseStorageError::StorageError) } - async fn get_response_chain( - &self, - response_id: &ResponseId, - max_depth: Option, - ) -> Result { - let mut chain = ResponseChain::new(); - let mut current_id = Some(response_id.clone()); - let mut visited = 0usize; - - while let Some(ref lookup_id) = current_id { - if let Some(limit) = max_depth { - if visited >= limit { - break; - } - } - - let fetched = self.get_response(lookup_id).await?; - match fetched { - Some(response) => { - current_id.clone_from(&response.previous_response_id); - chain.responses.push(response); - visited += 1; - } - None => break, - } - } - - chain.responses.reverse(); - Ok(chain) - } - async fn list_identifier_responses( &self, identifier: &str, limit: Option, ) -> Result, ResponseStorageError> { let identifier = identifier.to_string(); + let select_base = self.select_base.clone(); + let schema = self.store.schema.clone(); self.store .execute(move |conn| { + let s = &schema.responses; + let col_safety = s.col("safety_identifier"); + let col_created = s.col("created_at"); + let sql = if let Some(limit) = limit { format!( - "SELECT * FROM ({SELECT_BASE} WHERE safety_identifier = :1 ORDER BY created_at DESC) WHERE ROWNUM <= {limit}" + "SELECT * FROM ({select_base} WHERE {col_safety} = :1 \ + ORDER BY {col_created} DESC) WHERE ROWNUM <= {limit}" ) } else { - format!("{SELECT_BASE} WHERE safety_identifier = :1 ORDER BY created_at DESC") + format!("{select_base} WHERE {col_safety} = :1 ORDER BY {col_created} DESC") }; let mut stmt = conn.statement(&sql).build().map_err(map_oracle_error)?; @@ -1096,7 +1311,7 @@ impl ResponseStorage for OracleResponseStorage { for row in &mut rows { let row = row.map_err(map_oracle_error)?; - results.push(Self::build_response_from_row(&row)?); + results.push(Self::build_response_from_row(&row, &schema)?); } Ok(results) @@ -1110,11 +1325,17 @@ impl ResponseStorage for OracleResponseStorage { identifier: &str, ) -> Result { let identifier = identifier.to_string(); + let schema = self.store.schema.clone(); + let affected = self .store .execute(move |conn| { + let s = &schema.responses; + let table = s.qualified_table(schema.owner.as_deref()); + let col_safety = s.col("safety_identifier"); + conn.execute( - "DELETE FROM responses WHERE safety_identifier = :1", + &format!("DELETE FROM {table} WHERE {col_safety} = :1"), &[&identifier], ) .map_err(map_oracle_error) @@ -1130,12 +1351,18 @@ impl ResponseStorage for OracleResponseStorage { } } -// Helper functions for response parsing - -fn create_index_if_missing(conn: &Connection, index_name: &str, ddl: &str) -> Result<(), String> { +fn create_index_if_missing( + conn: &Connection, + table_upper: &str, + index_name: &str, + ddl: &str, +) -> Result<(), String> { let count: i64 = conn .query_row_as( - "SELECT COUNT(*) FROM user_indexes WHERE table_name = 'RESPONSES' AND index_name = :1", + &format!( + "SELECT COUNT(*) FROM user_indexes \ + WHERE table_name = '{table_upper}' AND index_name = :1" + ), &[&index_name], ) .map_err(map_oracle_error)?; @@ -1145,7 +1372,6 @@ fn create_index_if_missing(conn: &Connection, index_name: &str, ddl: &str) -> Re if let Some(db_err) = err.db_error() { // ORA-00955: name is already used by an existing object // ORA-01408: such column list already indexed - // Both errors indicate the index already exists (race condition) if db_err.code() != 955 && db_err.code() != 1408 { return Err(map_oracle_error(err)); } diff --git a/data_connector/src/postgres.rs b/data_connector/src/postgres.rs index b9b9d10a7a..550a526a29 100644 --- a/data_connector/src/postgres.rs +++ b/data_connector/src/postgres.rs @@ -6,7 +6,7 @@ //! 3. PostgresConversationItemStorage //! 4. PostgresResponseStorage -use std::str::FromStr; +use std::{str::FromStr, sync::Arc}; use async_trait::async_trait; use chrono::{DateTime, Utc}; @@ -15,23 +15,32 @@ use serde_json::Value; use tokio_postgres::{NoTls, Row}; use crate::{ - common::{parse_json_value, parse_metadata, parse_raw_response, parse_tool_calls}, + common::{ + build_response_select_base, parse_json_value, parse_metadata, parse_raw_response, + parse_tool_calls, + }, config::PostgresConfig, core::{ make_item_id, Conversation, ConversationId, ConversationItem, ConversationItemId, ConversationItemResult, ConversationItemStorage, ConversationItemStorageError, ConversationMetadata, ConversationResult, ConversationStorage, ConversationStorageError, - ListParams, NewConversation, NewConversationItem, ResponseChain, ResponseId, - ResponseResult, ResponseStorage, ResponseStorageError, SortOrder, StoredResponse, + ListParams, NewConversation, NewConversationItem, ResponseId, ResponseResult, + ResponseStorage, ResponseStorageError, SortOrder, StoredResponse, }, + schema::SchemaConfig, }; pub(crate) struct PostgresStore { pool: Pool, + pub(crate) schema: Arc, } impl PostgresStore { pub fn new(config: PostgresConfig) -> Result { + let schema = config.schema.clone().unwrap_or_default(); + schema.validate()?; + let schema = Arc::new(schema); + let pg_config = tokio_postgres::Config::from_str(config.db_url.as_str()) .map_err(|e| format!("Invalid PostgreSQL connection URL: {e}"))?; let mgr_config = ManagerConfig { @@ -43,7 +52,7 @@ impl PostgresStore { .build() .map_err(|e| format!("Failed to build PostgreSQL connection pool: {e}"))?; - Ok(Self { pool }) + Ok(Self { pool, schema }) } } @@ -51,6 +60,7 @@ impl Clone for PostgresStore { fn clone(&self) -> Self { Self { pool: self.pool.clone(), + schema: self.schema.clone(), } } } @@ -61,20 +71,27 @@ pub(super) struct PostgresConversationStorage { impl PostgresConversationStorage { pub async fn new(store: PostgresStore) -> Result { + let s = &store.schema.conversations; + let table = s.qualified_table(store.schema.owner.as_deref()); + + let col_defs = [ + format!("{} VARCHAR(64) PRIMARY KEY", s.col("id")), + format!("{} TIMESTAMPTZ", s.col("created_at")), + format!("{} JSON", s.col("metadata")), + ]; + + let ddl = format!( + "CREATE TABLE IF NOT EXISTS {table} ({});", + col_defs.join(", ") + ); + let client = store .pool .get() .await .map_err(|e| ConversationStorageError::StorageError(e.to_string()))?; client - .batch_execute( - " - CREATE TABLE IF NOT EXISTS conversations ( - id VARCHAR(64) PRIMARY KEY, - created_at TIMESTAMPTZ, - metadata JSON - );", - ) + .batch_execute(&ddl) .await .map_err(|e| ConversationStorageError::StorageError(e.to_string()))?; Ok(Self { store }) @@ -102,6 +119,17 @@ impl ConversationStorage for PostgresConversationStorage { .as_ref() .map(serde_json::to_string) .transpose()?; + + let s = &self.store.schema.conversations; + let table = s.qualified_table(self.store.schema.owner.as_deref()); + let col_id = s.col("id"); + let col_created = s.col("created_at"); + let col_meta = s.col("metadata"); + + let sql = format!( + "INSERT INTO {table} ({col_id}, {col_created}, {col_meta}) VALUES ($1, $2, $3)" + ); + let client = self .store .pool @@ -109,10 +137,7 @@ impl ConversationStorage for PostgresConversationStorage { .await .map_err(|e| ConversationStorageError::StorageError(e.to_string()))?; client - .execute( - "INSERT INTO conversations (id, created_at, metadata) VALUES ($1, $2, $3)", - &[&id_str, &created_at, &metadata_json], - ) + .execute(&sql, &[&id_str, &created_at, &metadata_json]) .await .map_err(|e| ConversationStorageError::StorageError(e.to_string()))?; Ok(conversation) @@ -122,6 +147,15 @@ impl ConversationStorage for PostgresConversationStorage { &self, id: &ConversationId, ) -> Result, ConversationStorageError> { + let s = &self.store.schema.conversations; + let table = s.qualified_table(self.store.schema.owner.as_deref()); + let col_id = s.col("id"); + let col_created = s.col("created_at"); + let col_meta = s.col("metadata"); + + let sql = + format!("SELECT {col_id}, {col_created}, {col_meta} FROM {table} WHERE {col_id} = $1"); + let client = self .store .pool @@ -129,19 +163,16 @@ impl ConversationStorage for PostgresConversationStorage { .await .map_err(|e| ConversationStorageError::StorageError(e.to_string()))?; let rows = client - .query( - "SELECT id, created_at, metadata FROM conversations WHERE id = $1", - &[&id.0.as_str()], - ) + .query(&sql, &[&id.0.as_str()]) .await .map_err(|e| ConversationStorageError::StorageError(e.to_string()))?; if rows.is_empty() { return Ok(None); } let row = &rows[0]; - let id_str: String = row.get(0); - let created_at: DateTime = row.get(1); - let metadata_json: Option = row.get(2); + let id_str: String = row.get(col_id); + let created_at: DateTime = row.get(col_created); + let metadata_json: Option = row.get(col_meta); let metadata = Self::parse_metadata(metadata_json)?; Ok(Some(Conversation::with_parts( ConversationId(id_str), @@ -156,6 +187,17 @@ impl ConversationStorage for PostgresConversationStorage { metadata: Option, ) -> Result, ConversationStorageError> { let metadata_json = metadata.as_ref().map(serde_json::to_string).transpose()?; + + let s = &self.store.schema.conversations; + let table = s.qualified_table(self.store.schema.owner.as_deref()); + let col_id = s.col("id"); + let col_meta = s.col("metadata"); + let col_created = s.col("created_at"); + + let sql = format!( + "UPDATE {table} SET {col_meta} = $1 WHERE {col_id} = $2 RETURNING {col_created}" + ); + let client = self .store .pool @@ -163,17 +205,14 @@ impl ConversationStorage for PostgresConversationStorage { .await .map_err(|e| ConversationStorageError::StorageError(e.to_string()))?; let rows = client - .query( - "UPDATE conversations SET metadata = $1 WHERE id = $2 RETURNING created_at", - &[&metadata_json, &id.0.as_str()], - ) + .query(&sql, &[&metadata_json, &id.0.as_str()]) .await .map_err(|e| ConversationStorageError::StorageError(e.to_string()))?; if rows.is_empty() { return Ok(None); } let row = &rows[0]; - let created_at: DateTime = row.get(0); + let created_at: DateTime = row.get(col_created); Ok(Some(Conversation::with_parts( ConversationId(id.0.clone()), created_at, @@ -182,6 +221,10 @@ impl ConversationStorage for PostgresConversationStorage { } async fn delete_conversation(&self, id: &ConversationId) -> ConversationResult { + let s = &self.store.schema.conversations; + let table = s.qualified_table(self.store.schema.owner.as_deref()); + let col_id = s.col("id"); + let client = self .store .pool @@ -189,7 +232,10 @@ impl ConversationStorage for PostgresConversationStorage { .await .map_err(|e| ConversationStorageError::StorageError(e.to_string()))?; let rows_deleted = client - .execute("DELETE FROM conversations WHERE id = $1", &[&id.0.as_str()]) + .execute( + &format!("DELETE FROM {table} WHERE {col_id} = $1"), + &[&id.0.as_str()], + ) .await .map_err(|e| ConversationStorageError::StorageError(e.to_string()))?; Ok(rows_deleted > 0) @@ -202,38 +248,52 @@ pub(super) struct PostgresConversationItemStorage { impl PostgresConversationItemStorage { pub async fn new(store: PostgresStore) -> Result { + let schema = &store.schema; + let si = &schema.conversation_items; + let sl = &schema.conversation_item_links; + let items_table = si.qualified_table(schema.owner.as_deref()); + let links_table = sl.qualified_table(schema.owner.as_deref()); + + // ── conversation_items DDL ── + let item_col_defs = [ + format!("{} VARCHAR(64) PRIMARY KEY", si.col("id")), + format!("{} VARCHAR(64)", si.col("response_id")), + format!("{} VARCHAR(32) NOT NULL", si.col("item_type")), + format!("{} VARCHAR(32)", si.col("role")), + format!("{} JSON", si.col("content")), + format!("{} VARCHAR(32)", si.col("status")), + format!("{} TIMESTAMPTZ", si.col("created_at")), + ]; + + // ── conversation_item_links DDL ── + let col_conv_id = sl.col("conversation_id"); + let col_item_id = sl.col("item_id"); + let col_added_at = sl.col("added_at"); + + let mut link_col_defs = vec![ + format!("{col_conv_id} VARCHAR(64)"), + format!("{col_item_id} VARCHAR(64) NOT NULL"), + format!("{col_added_at} TIMESTAMPTZ"), + ]; + link_col_defs.push(format!( + "CONSTRAINT pk_conv_item_link PRIMARY KEY ({col_conv_id}, {col_item_id})" + )); + + let ddl = format!( + "CREATE TABLE IF NOT EXISTS {items_table} ({});\n\ + CREATE TABLE IF NOT EXISTS {links_table} ({});\n\ + CREATE INDEX IF NOT EXISTS conv_item_links_conv_idx ON {links_table} ({col_conv_id}, {col_added_at});", + item_col_defs.join(", "), + link_col_defs.join(", "), + ); + let client = store .pool .get() .await .map_err(|e| ConversationItemStorageError::StorageError(e.to_string()))?; client - .batch_execute( - " - CREATE TABLE IF NOT EXISTS conversation_items ( - id VARCHAR(64) PRIMARY KEY, - response_id VARCHAR(64), - item_type VARCHAR(32) NOT NULL, - role VARCHAR(32), - content JSON, - status VARCHAR(32), - created_at TIMESTAMPTZ - );", - ) - .await - .map_err(|e| ConversationItemStorageError::StorageError(e.to_string()))?; - - client - .batch_execute( - " - CREATE TABLE IF NOT EXISTS conversation_item_links ( - conversation_id VARCHAR(64), - item_id VARCHAR(64) NOT NULL, - added_at TIMESTAMPTZ, - CONSTRAINT pk_conv_item_link PRIMARY KEY (conversation_id, item_id) - ); - CREATE INDEX IF NOT EXISTS conv_item_links_conv_idx ON conversation_item_links (conversation_id, added_at);", - ) + .batch_execute(&ddl) .await .map_err(|e| ConversationItemStorageError::StorageError(e.to_string()))?; Ok(Self { store }) @@ -258,14 +318,40 @@ impl ConversationItemStorage for PostgresConversationItemStorage { let created_at = Utc::now(); let content_json = serde_json::to_string(&content)?; + let si = &self.store.schema.conversation_items; + let table = si.qualified_table(self.store.schema.owner.as_deref()); + + let col_id = si.col("id"); + let col_response_id = si.col("response_id"); + let col_item_type = si.col("item_type"); + let col_role = si.col("role"); + let col_content = si.col("content"); + let col_status = si.col("status"); + let col_created_at = si.col("created_at"); + + let sql = format!( + "INSERT INTO {table} ({col_id}, {col_response_id}, {col_item_type}, {col_role}, {col_content}, {col_status}, {col_created_at}) \ + VALUES ($1, $2, $3, $4, $5, $6, $7)" + ); + + let params: Vec<&(dyn tokio_postgres::types::ToSql + Sync)> = vec![ + &id.0, + &response_id, + &item_type, + &role, + &content_json, + &status, + &created_at, + ]; + let client = self .store .pool .get() .await .map_err(|e| ConversationItemStorageError::StorageError(e.to_string()))?; - client.execute("INSERT INTO conversation_items (id, response_id, item_type, role, content, status, created_at) VALUES ($1, $2, $3, $4, $5, $6, $7)", - &[&id.0.as_str(), &response_id, &item_type, &role, &content_json, &status, &created_at]) + client + .execute(&sql, ¶ms) .await .map_err(|e| ConversationItemStorageError::StorageError(e.to_string()))?; Ok(ConversationItem { @@ -285,14 +371,27 @@ impl ConversationItemStorage for PostgresConversationItemStorage { item_id: &ConversationItemId, added_at: DateTime, ) -> ConversationItemResult<()> { + let sl = &self.store.schema.conversation_item_links; + let table = sl.qualified_table(self.store.schema.owner.as_deref()); + let col_conv = sl.col("conversation_id"); + let col_item = sl.col("item_id"); + let col_added = sl.col("added_at"); + + let sql = format!( + "INSERT INTO {table} ({col_conv}, {col_item}, {col_added}) VALUES ($1, $2, $3)" + ); + let client = self .store .pool .get() .await .map_err(|e| ConversationItemStorageError::StorageError(e.to_string()))?; - client.execute("INSERT INTO conversation_item_links (conversation_id, item_id, added_at) VALUES ($1, $2, $3)", - &[&conversation_id.0.as_str(), &item_id.0.as_str(), &added_at]) + client + .execute( + &sql, + &[&conversation_id.0.as_str(), &item_id.0.as_str(), &added_at], + ) .await .map_err(|e| ConversationItemStorageError::StorageError(e.to_string()))?; Ok(()) @@ -303,6 +402,17 @@ impl ConversationItemStorage for PostgresConversationItemStorage { conversation_id: &ConversationId, params: ListParams, ) -> ConversationItemResult> { + let schema = &self.store.schema; + let si = &schema.conversation_items; + let sl = &schema.conversation_item_links; + let items_table = si.qualified_table(schema.owner.as_deref()); + let links_table = sl.qualified_table(schema.owner.as_deref()); + + let l_conv_id = sl.col("conversation_id"); + let l_item_id = sl.col("item_id"); + let l_added_at = sl.col("added_at"); + let i_id = si.col("id"); + let cid = conversation_id.0.as_str(); let limit: i64 = params.limit as i64; let order_desc = matches!(params.order, SortOrder::Desc); @@ -314,11 +424,11 @@ impl ConversationItemStorage for PostgresConversationItemStorage { .get() .await .map_err(|e| ConversationItemStorageError::StorageError(e.to_string()))?; + let cursor_sql = format!( + "SELECT {l_added_at} FROM {links_table} WHERE {l_conv_id} = $1 AND {l_item_id} = $2" + ); let rows = client - .query( - "SELECT added_at FROM conversation_item_links WHERE conversation_id = $1 AND item_id = $2", - &[&cid, &aid.as_str()], - ) + .query(&cursor_sql, &[&cid, &aid.as_str()]) .await .map_err(|e| ConversationItemStorageError::StorageError(e.to_string()))?; if rows.is_empty() { @@ -332,32 +442,50 @@ impl ConversationItemStorage for PostgresConversationItemStorage { None }; - let mut sql = String::from( - "SELECT i.id, i.response_id, i.item_type, i.role, i.content, i.status, i.created_at \ - FROM conversation_item_links l \ - JOIN conversation_items i ON i.id = l.item_id \ - WHERE l.conversation_id = $1", + // Build select columns from items table (prefixed with i.) + let mut select_cols = Vec::new(); + for field in &[ + "id", + "response_id", + "item_type", + "role", + "content", + "status", + "created_at", + ] { + select_cols.push(format!("i.{}", si.col(field))); + } + + let mut sql = format!( + "SELECT {} FROM {links_table} l JOIN {items_table} i ON i.{i_id} = l.{l_item_id} \ + WHERE l.{l_conv_id} = $1", + select_cols.join(", "), ); - // If cursor provided, append predicate using $2/$3 + if let Some((_ts, _iid)) = &after_key { if order_desc { - sql.push_str(" AND (l.added_at < $2 OR (l.added_at = $2 AND l.item_id < $3))"); + sql.push_str(&format!( + " AND (l.{l_added_at} < $2 OR (l.{l_added_at} = $2 AND l.{l_item_id} < $3))" + )); } else { - sql.push_str(" AND (l.added_at > $2 OR (l.added_at = $2 AND l.item_id > $3))"); + sql.push_str(&format!( + " AND (l.{l_added_at} > $2 OR (l.{l_added_at} = $2 AND l.{l_item_id} > $3))" + )); } } - // Order and limit if order_desc { - sql.push_str(" ORDER BY l.added_at DESC, l.item_id DESC"); + sql.push_str(&format!( + " ORDER BY l.{l_added_at} DESC, l.{l_item_id} DESC" + )); } else { - sql.push_str(" ORDER BY l.added_at ASC, l.item_id ASC"); + sql.push_str(&format!(" ORDER BY l.{l_added_at} ASC, l.{l_item_id} ASC")); } - // PostgreSQL LIMIT if after_key.is_some() { sql.push_str(" LIMIT $4"); } else { sql.push_str(" LIMIT $2"); } + let client = self .store .pool @@ -375,52 +503,68 @@ impl ConversationItemStorage for PostgresConversationItemStorage { .await .map_err(|e| ConversationItemStorageError::StorageError(e.to_string()))? }; + + let col_id = si.col("id"); + let col_resp_id = si.col("response_id"); + let col_item_type = si.col("item_type"); + let col_role = si.col("role"); + let col_content = si.col("content"); + let col_status = si.col("status"); + let col_created_at = si.col("created_at"); + let mut out = Vec::new(); for row in rows { - let id = row.get(0); - let resp_id: Option = row.get(1); - let item_type: String = row.get(2); - let role: Option = row.get(3); - let content_raw: Option = row.get(4); - let status: Option = row.get(5); - let created_at: DateTime = row.get(6); - out.push(( - id, - resp_id, + let id: String = row.get(col_id); + let resp_id: Option = row.get(col_resp_id); + let item_type: String = row.get(col_item_type); + let role: Option = row.get(col_role); + let content_raw: Option = row.get(col_content); + let status: Option = row.get(col_status); + let created_at: DateTime = row.get(col_created_at); + + let content = match content_raw { + Some(s) => serde_json::from_str(&s).map_err(ConversationItemStorageError::from)?, + None => Value::Null, + }; + out.push(ConversationItem { + id: ConversationItemId(id), + response_id: resp_id, item_type, role, - content_raw, + content, status, created_at, - )); + }); } - out.into_iter() - .map( - |(id, resp_id, item_type, role, content_raw, status, created_at)| { - let content = match content_raw { - Some(s) => { - serde_json::from_str(&s).map_err(ConversationItemStorageError::from)? - } - None => Value::Null, - }; - Ok(ConversationItem { - id: ConversationItemId(id), - response_id: resp_id, - item_type, - role, - content, - status, - created_at, - }) - }, - ) - .collect() + Ok(out) } async fn get_item( &self, item_id: &ConversationItemId, ) -> Result, ConversationItemStorageError> { + let si = &self.store.schema.conversation_items; + let table = si.qualified_table(self.store.schema.owner.as_deref()); + let col_id = si.col("id"); + + let mut select_cols = Vec::new(); + for field in &[ + "id", + "response_id", + "item_type", + "role", + "content", + "status", + "created_at", + ] { + select_cols.push(si.col(field).to_string()); + } + + let sql = format!( + "SELECT {} FROM {table} WHERE {col_id} = $1", + select_cols.join(", "), + ); + let client = self .store .pool @@ -428,23 +572,20 @@ impl ConversationItemStorage for PostgresConversationItemStorage { .await .map_err(|e| ConversationItemStorageError::StorageError(e.to_string()))?; let rows = client - .query( - "SELECT id, response_id, item_type, role, content, status, created_at FROM conversation_items WHERE id = $1", - &[&item_id.0.as_str()], - ) + .query(&sql, &[&item_id.0.as_str()]) .await .map_err(|e| ConversationItemStorageError::StorageError(e.to_string()))?; if rows.is_empty() { Ok(None) } else { let row = &rows[0]; - let id: String = row.get(0); - let response_id: Option = row.get(1); - let item_type: String = row.get(2); - let role: Option = row.get(3); - let content_raw: Option = row.get(4); - let status: Option = row.get(5); - let created_at: DateTime = row.get(6); + let id: String = row.get(si.col("id")); + let response_id: Option = row.get(si.col("response_id")); + let item_type: String = row.get(si.col("item_type")); + let role: Option = row.get(si.col("role")); + let content_raw: Option = row.get(si.col("content")); + let status: Option = row.get(si.col("status")); + let created_at: DateTime = row.get(si.col("created_at")); let content = match content_raw { Some(s) => serde_json::from_str(&s) @@ -469,6 +610,13 @@ impl ConversationItemStorage for PostgresConversationItemStorage { conversation_id: &ConversationId, item_id: &ConversationItemId, ) -> ConversationItemResult { + let sl = &self.store.schema.conversation_item_links; + let table = sl.qualified_table(self.store.schema.owner.as_deref()); + let col_conv = sl.col("conversation_id"); + let col_item = sl.col("item_id"); + + let sql = format!("SELECT COUNT(*) FROM {table} WHERE {col_conv} = $1 AND {col_item} = $2"); + let client = self .store .pool @@ -476,10 +624,7 @@ impl ConversationItemStorage for PostgresConversationItemStorage { .await .map_err(|e| ConversationItemStorageError::StorageError(e.to_string()))?; let row = client - .query_one( - "SELECT COUNT(*) FROM conversation_item_links WHERE conversation_id = $1 AND item_id = $2", - &[&conversation_id.0.as_str(), &item_id.0.as_str()], - ) + .query_one(&sql, &[&conversation_id.0.as_str(), &item_id.0.as_str()]) .await .map_err(|e| ConversationItemStorageError::StorageError(e.to_string()))?; let count: i64 = row.get(0); @@ -491,6 +636,13 @@ impl ConversationItemStorage for PostgresConversationItemStorage { conversation_id: &ConversationId, item_id: &ConversationItemId, ) -> ConversationItemResult<()> { + let sl = &self.store.schema.conversation_item_links; + let table = sl.qualified_table(self.store.schema.owner.as_deref()); + let col_conv = sl.col("conversation_id"); + let col_item = sl.col("item_id"); + + let sql = format!("DELETE FROM {table} WHERE {col_conv} = $1 AND {col_item} = $2"); + let client = self .store .pool @@ -498,10 +650,7 @@ impl ConversationItemStorage for PostgresConversationItemStorage { .await .map_err(|e| ConversationItemStorageError::StorageError(e.to_string()))?; client - .execute( - "DELETE FROM conversation_item_links WHERE conversation_id = $1 AND item_id = $2", - &[&conversation_id.0.as_str(), &item_id.0.as_str()], - ) + .execute(&sql, &[&conversation_id.0.as_str(), &item_id.0.as_str()]) .await .map_err(|e| ConversationItemStorageError::StorageError(e.to_string()))?; Ok(()) @@ -510,52 +659,71 @@ impl ConversationItemStorage for PostgresConversationItemStorage { pub(super) struct PostgresResponseStorage { store: PostgresStore, + select_base: String, } impl PostgresResponseStorage { pub async fn new(store: PostgresStore) -> Result { + let schema = &store.schema; + let s = &schema.responses; + let table = s.qualified_table(schema.owner.as_deref()); + + let col_defs = vec![ + format!("{} VARCHAR(64) PRIMARY KEY", s.col("id")), + format!("{} VARCHAR(64)", s.col("conversation_id")), + format!("{} VARCHAR(64)", s.col("previous_response_id")), + format!("{} JSON", s.col("input")), + format!("{} TEXT", s.col("instructions")), + format!("{} JSON", s.col("output")), + format!("{} JSON", s.col("tool_calls")), + format!("{} JSON", s.col("metadata")), + format!("{} TIMESTAMPTZ", s.col("created_at")), + format!("{} VARCHAR(128)", s.col("safety_identifier")), + format!("{} VARCHAR(128)", s.col("model")), + format!("{} JSON", s.col("raw_response")), + ]; + + let mut ddl = format!( + "CREATE TABLE IF NOT EXISTS {table} ({});", + col_defs.join(", "), + ); + ddl.push_str(&format!( + "\nCREATE INDEX IF NOT EXISTS responses_safety_idx ON {table} ({});", + s.col("safety_identifier") + )); + let client = store .pool .get() .await .map_err(|e| ResponseStorageError::StorageError(e.to_string()))?; client - .batch_execute( - " - CREATE TABLE IF NOT EXISTS responses ( - id VARCHAR(64) PRIMARY KEY, - conversation_id VARCHAR(64), - previous_response_id VARCHAR(64), - input JSON, - instructions TEXT, - output JSON, - tool_calls JSON, - metadata JSON, - created_at TIMESTAMPTZ, - safety_identifier VARCHAR(128), - model VARCHAR(128), - raw_response JSON - ); - CREATE INDEX IF NOT EXISTS responses_safety_idx ON responses (safety_identifier);", - ) + .batch_execute(&ddl) .await .map_err(|e| ResponseStorageError::StorageError(e.to_string()))?; - Ok(Self { store }) + + let select_base = build_response_select_base(&store.schema); + Ok(Self { store, select_base }) } - pub fn build_response_from_row(row: &Row) -> Result { - let id: String = row.get("id"); - let conversation_id: Option = row.get("conversation_id"); - let previous: Option = row.get("previous_response_id"); - let input_json: Option = row.get("input"); - let instructions: Option = row.get("instructions"); - let output_json: Option = row.get("output"); - let tool_calls_json: Option = row.get("tool_calls"); - let metadata_json: Option = row.get("metadata"); - let created_at: DateTime = row.get("created_at"); - let safety_identifier: Option = row.get("safety_identifier"); - let model: Option = row.get("model"); - let raw_response_json: Option = row.get("raw_response"); + pub fn build_response_from_row( + row: &Row, + schema: &SchemaConfig, + ) -> Result { + let s = &schema.responses; + + let id: String = row.get(s.col("id")); + let conversation_id: Option = row.get(s.col("conversation_id")); + let previous: Option = row.get(s.col("previous_response_id")); + let input_json: Option = row.get(s.col("input")); + let instructions: Option = row.get(s.col("instructions")); + let output_json: Option = row.get(s.col("output")); + let tool_calls_json: Option = row.get(s.col("tool_calls")); + let metadata_json: Option = row.get(s.col("metadata")); + let created_at: DateTime = row.get(s.col("created_at")); + let safety_identifier: Option = row.get(s.col("safety_identifier")); + let model: Option = row.get(s.col("model")); + let raw_response_json: Option = row.get(s.col("raw_response")); let previous_response_id = previous.map(ResponseId); let tool_calls = parse_tool_calls(tool_calls_json)?; @@ -604,30 +772,54 @@ impl ResponseStorage for PostgresResponseStorage { let previous_id = previous_response_id.map(|r| r.0); let tool_calls_value = serde_json::to_value(&tool_calls)?; let metadata_value = serde_json::to_value(&metadata)?; + + let s = &self.store.schema.responses; + let table = s.qualified_table(self.store.schema.owner.as_deref()); + + let col_id = s.col("id"); + let col_prev = s.col("previous_response_id"); + let col_input = s.col("input"); + let col_instructions = s.col("instructions"); + let col_output = s.col("output"); + let col_tool_calls = s.col("tool_calls"); + let col_metadata = s.col("metadata"); + let col_created_at = s.col("created_at"); + let col_safety = s.col("safety_identifier"); + let col_model = s.col("model"); + let col_conv = s.col("conversation_id"); + let col_raw = s.col("raw_response"); + + let sql = format!( + "INSERT INTO {table} ({col_id}, {col_prev}, {col_input}, {col_instructions}, {col_output}, \ + {col_tool_calls}, {col_metadata}, {col_created_at}, {col_safety}, {col_model}, {col_conv}, {col_raw}) \ + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12)" + ); + + let params: Vec<&(dyn tokio_postgres::types::ToSql + Sync)> = vec![ + &response_id.0, + &previous_id, + &input, + &instructions, + &output, + &tool_calls_value, + &metadata_value, + &created_at, + &safety_identifier, + &model, + &conversation_id, + &raw_response, + ]; + let client = self .store .pool .get() .await .map_err(|e| ResponseStorageError::StorageError(e.to_string()))?; - let insert_count = client.execute( - "INSERT INTO responses (id, previous_response_id, input, instructions, output, \ - tool_calls, metadata, created_at, safety_identifier, model, conversation_id, raw_response) \ - VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12)", - &[ - &response_id.0.as_str(), - &previous_id, - &input, - &instructions, - &output, - &tool_calls_value, - &metadata_value, - &created_at, - &safety_identifier, - &model, - &conversation_id, - &raw_response, - ]).await.map_err(|e| ResponseStorageError::StorageError(e.to_string()))?; + let insert_count = client + .execute(&sql, ¶ms) + .await + .map_err(|e| ResponseStorageError::StorageError(e.to_string()))?; tracing::debug!(rows_affected = insert_count, "Response stored in Postgres"); Ok(response_id) } @@ -636,6 +828,9 @@ impl ResponseStorage for PostgresResponseStorage { &self, response_id: &ResponseId, ) -> Result, ResponseStorageError> { + let col_id = self.store.schema.responses.col("id"); + let sql = format!("{} WHERE {col_id} = $1", self.select_base); + let client = self .store .pool @@ -643,21 +838,22 @@ impl ResponseStorage for PostgresResponseStorage { .await .map_err(|e| ResponseStorageError::StorageError(e.to_string()))?; let rows = client - .query( - "SELECT * FROM responses WHERE id = $1", - &[&response_id.0.as_str()], - ) + .query(&sql, &[&response_id.0.as_str()]) .await .map_err(|e| ResponseStorageError::StorageError(e.to_string()))?; if rows.is_empty() { return Ok(None); } - Self::build_response_from_row(&rows[0]) + Self::build_response_from_row(&rows[0], &self.store.schema) .map(Some) .map_err(|err| ResponseStorageError::StorageError(err.to_string())) } async fn delete_response(&self, response_id: &ResponseId) -> ResponseResult<()> { + let s = &self.store.schema.responses; + let table = s.qualified_table(self.store.schema.owner.as_deref()); + let col_id = s.col("id"); + let client = self .store .pool @@ -666,7 +862,7 @@ impl ResponseStorage for PostgresResponseStorage { .map_err(|e| ResponseStorageError::StorageError(e.to_string()))?; client .execute( - "DELETE FROM responses WHERE id = $1", + &format!("DELETE FROM {table} WHERE {col_id} = $1"), &[&response_id.0.as_str()], ) .await @@ -674,42 +870,15 @@ impl ResponseStorage for PostgresResponseStorage { Ok(()) } - async fn get_response_chain( - &self, - response_id: &ResponseId, - max_depth: Option, - ) -> ResponseResult { - let mut chain = ResponseChain::new(); - let mut current_id = Some(response_id.clone()); - let mut visited = 0usize; - - while let Some(ref lookup_id) = current_id { - if let Some(limit) = max_depth { - if visited >= limit { - break; - } - } - - let fetched = self.get_response(lookup_id).await?; - match fetched { - Some(response) => { - current_id.clone_from(&response.previous_response_id); - chain.responses.push(response); - visited += 1; - } - None => break, - } - } - - chain.responses.reverse(); - Ok(chain) - } - async fn list_identifier_responses( &self, identifier: &str, limit: Option, ) -> ResponseResult> { + let s = &self.store.schema.responses; + let col_safety = s.col("safety_identifier"); + let col_created = s.col("created_at"); + let client = self .store .pool @@ -718,27 +887,30 @@ impl ResponseStorage for PostgresResponseStorage { .map_err(|e| ResponseStorageError::StorageError(e.to_string()))?; let rows = if let Some(l) = limit { let l_i64: i64 = l as i64; + let sql = format!( + "{} WHERE {col_safety} = $1 ORDER BY {col_created} DESC LIMIT $2", + self.select_base, + ); client - .query( - "SELECT * FROM responses WHERE safety_identifier = $1 ORDER BY created_at DESC LIMIT $2", - &[&identifier, &l_i64], - ) + .query(&sql, &[&identifier, &l_i64]) .await .map_err(|e| ResponseStorageError::StorageError(e.to_string()))? } else { + let sql = format!( + "{} WHERE {col_safety} = $1 ORDER BY {col_created} DESC", + self.select_base, + ); client - .query( - "SELECT * FROM responses WHERE safety_identifier = $1 ORDER BY created_at DESC", - &[&identifier], - ) + .query(&sql, &[&identifier]) .await .map_err(|e| ResponseStorageError::StorageError(e.to_string()))? }; + let schema = &self.store.schema; let mut out = Vec::with_capacity(rows.len()); for row in rows { - let resp = - Self::build_response_from_row(&row).map_err(ResponseStorageError::StorageError)?; + let resp = Self::build_response_from_row(&row, schema) + .map_err(ResponseStorageError::StorageError)?; out.push(resp); } @@ -746,6 +918,12 @@ impl ResponseStorage for PostgresResponseStorage { } async fn delete_identifier_responses(&self, identifier: &str) -> ResponseResult { + let s = &self.store.schema.responses; + let table = s.qualified_table(self.store.schema.owner.as_deref()); + let col_safety = s.col("safety_identifier"); + + let sql = format!("DELETE FROM {table} WHERE {col_safety} = $1"); + let client = self .store .pool @@ -753,10 +931,7 @@ impl ResponseStorage for PostgresResponseStorage { .await .map_err(|e| ResponseStorageError::StorageError(e.to_string()))?; let rows_deleted = client - .execute( - "DELETE FROM responses WHERE safety_identifier = $1", - &[&identifier], - ) + .execute(&sql, &[&identifier]) .await .map_err(|e| ResponseStorageError::StorageError(e.to_string()))?; Ok(rows_deleted as usize) diff --git a/data_connector/src/redis.rs b/data_connector/src/redis.rs index 5a93739e58..de91f54ef5 100644 --- a/data_connector/src/redis.rs +++ b/data_connector/src/redis.rs @@ -6,6 +6,8 @@ //! 3. RedisConversationItemStorage //! 4. RedisResponseStorage +use std::sync::Arc; + use async_trait::async_trait; use chrono::{DateTime, Utc}; use deadpool_redis::{Config, Pool, Runtime}; @@ -19,18 +21,24 @@ use crate::{ make_item_id, Conversation, ConversationId, ConversationItem, ConversationItemId, ConversationItemResult, ConversationItemStorage, ConversationItemStorageError, ConversationMetadata, ConversationResult, ConversationStorage, ConversationStorageError, - ListParams, NewConversation, NewConversationItem, ResponseChain, ResponseId, - ResponseResult, ResponseStorage, ResponseStorageError, SortOrder, StoredResponse, + ListParams, NewConversation, NewConversationItem, ResponseId, ResponseResult, + ResponseStorage, ResponseStorageError, SortOrder, StoredResponse, }, + schema::SchemaConfig, }; pub(crate) struct RedisStore { pool: Pool, retention_days: Option, + pub(crate) schema: Arc, } impl RedisStore { pub fn new(config: RedisConfig) -> Result { + let schema = config.schema.clone().unwrap_or_default(); + schema.validate()?; + let schema = Arc::new(schema); + let mut cfg = Config::from_url(config.url); cfg.pool = Some(deadpool_redis::PoolConfig::new(config.pool_max)); let pool = cfg @@ -39,6 +47,7 @@ impl RedisStore { Ok(Self { pool, retention_days: config.retention_days, + schema, }) } } @@ -48,6 +57,7 @@ impl Clone for RedisStore { Self { pool: self.pool.clone(), retention_days: self.retention_days, + schema: self.schema.clone(), } } } @@ -61,8 +71,11 @@ impl RedisConversationStorage { Self { store } } - fn conversation_key(id: &str) -> String { - format!("conversation:{id}") + fn conversation_key(&self, id: &str) -> String { + match &self.store.schema.owner { + Some(owner) => format!("{owner}:conversation:{id}"), + None => format!("conversation:{id}"), + } } fn parse_metadata( @@ -88,22 +101,23 @@ impl ConversationStorage for RedisConversationStorage { .map(serde_json::to_string) .transpose()?; + let s = &self.store.schema.conversations; + let mut conn = self .store .pool .get() .await .map_err(|e| ConversationStorageError::StorageError(e.to_string()))?; - let key = Self::conversation_key(id_str); + let key = self.conversation_key(id_str); let mut pipe = redis::pipe(); - pipe.hset(&key, "id", id_str); - pipe.hset(&key, "created_at", created_at.to_rfc3339()); + pipe.hset(&key, s.col("id"), id_str); + pipe.hset(&key, s.col("created_at"), created_at.to_rfc3339()); if let Some(meta) = metadata_json { - pipe.hset(&key, "metadata", meta); + pipe.hset(&key, s.col("metadata"), meta); } - // Expire after configured retention days (optional) if let Some(days) = self.store.retention_days { pipe.expire(&key, (days * 24 * 60 * 60) as i64); } @@ -119,8 +133,12 @@ impl ConversationStorage for RedisConversationStorage { &self, id: &ConversationId, ) -> Result, ConversationStorageError> { + let s = &self.store.schema.conversations; + let col_created = s.col("created_at"); + let col_meta = s.col("metadata"); + let id_str = id.0.as_str(); - let key = Self::conversation_key(id_str); + let key = self.conversation_key(id_str); let mut conn = self .store .pool @@ -137,8 +155,8 @@ impl ConversationStorage for RedisConversationStorage { } let (created_at_str, metadata_json): (String, Option) = redis::pipe() - .hget(&key, "created_at") - .hget(&key, "metadata") + .hget(&key, col_created) + .hget(&key, col_meta) .query_async(&mut conn) .await .map_err(|e| ConversationStorageError::StorageError(e.to_string()))?; @@ -161,8 +179,12 @@ impl ConversationStorage for RedisConversationStorage { id: &ConversationId, metadata: Option, ) -> Result, ConversationStorageError> { + let s = &self.store.schema.conversations; + let col_meta = s.col("metadata"); + let col_created = s.col("created_at"); + let id_str = id.0.as_str(); - let key = Self::conversation_key(id_str); + let key = self.conversation_key(id_str); let mut conn = self .store .pool @@ -181,18 +203,17 @@ impl ConversationStorage for RedisConversationStorage { let metadata_json = metadata.as_ref().map(serde_json::to_string).transpose()?; if let Some(meta) = metadata_json { - conn.hset::<_, _, _, ()>(&key, "metadata", meta) + conn.hset::<_, _, _, ()>(&key, col_meta, meta) .await .map_err(|e| ConversationStorageError::StorageError(e.to_string()))?; } else { - conn.hdel::<_, _, ()>(&key, "metadata") + conn.hdel::<_, _, ()>(&key, col_meta) .await .map_err(|e| ConversationStorageError::StorageError(e.to_string()))?; } - // We need to fetch created_at to return the full object let created_at_str: String = conn - .hget(&key, "created_at") + .hget(&key, col_created) .await .map_err(|e| ConversationStorageError::StorageError(e.to_string()))?; let created_at = DateTime::parse_from_rfc3339(&created_at_str) @@ -208,8 +229,7 @@ impl ConversationStorage for RedisConversationStorage { async fn delete_conversation(&self, id: &ConversationId) -> ConversationResult { let id_str = id.0.as_str(); - let key = Self::conversation_key(id_str); - // Also delete the items list for this conversation + let key = self.conversation_key(id_str); let items_key = format!("{key}:items"); let mut conn = self @@ -239,55 +259,70 @@ impl RedisConversationItemStorage { Self { store } } - fn item_key(id: &str) -> String { - format!("item:{id}") + fn item_key(&self, id: &str) -> String { + match &self.store.schema.owner { + Some(owner) => format!("{owner}:item:{id}"), + None => format!("item:{id}"), + } } - fn conv_items_key(conv_id: &str) -> String { - format!("conversation:{conv_id}:items") + fn conv_items_key(&self, conv_id: &str) -> String { + match &self.store.schema.owner { + Some(owner) => format!("{owner}:conversation:{conv_id}:items"), + None => format!("conversation:{conv_id}:items"), + } } /// Parse a Redis hash map into a `ConversationItem`, returning errors for /// corrupted data instead of silently substituting defaults. fn build_item_from_map( + &self, map: &std::collections::HashMap, fallback_id: &str, ) -> Result { + let si = &self.store.schema.conversation_items; + + let col_id = si.col("id"); let id = ConversationItemId( - map.get("id") + map.get(col_id) .cloned() .unwrap_or_else(|| fallback_id.to_string()), ); - let response_id = map.get("response_id").cloned(); + + let response_id = map.get(si.col("response_id")).cloned(); + + let col_item_type = si.col("item_type"); let item_type = map - .get("item_type") + .get(col_item_type) .filter(|s| !s.is_empty()) .cloned() .ok_or_else(|| { ConversationItemStorageError::StorageError(format!( - "item {fallback_id} missing item_type" + "item {fallback_id} missing {col_item_type}" )) })?; - let role = map.get("role").cloned(); - let status = map.get("status").cloned(); - let content = match map.get("content") { + let role = map.get(si.col("role")).cloned(); + let status = map.get(si.col("status")).cloned(); + + let content = match map.get(si.col("content")) { Some(s) => { serde_json::from_str(s).map_err(ConversationItemStorageError::SerializationError)? } None => Value::Null, }; - let created_at_str = map.get("created_at").ok_or_else(|| { + let col_created = si.col("created_at"); + let created_at_str = map.get(col_created).ok_or_else(|| { ConversationItemStorageError::StorageError(format!( - "item {fallback_id} missing created_at" + "item {fallback_id} missing {col_created}" )) })?; let created_at = DateTime::parse_from_rfc3339(created_at_str) .map(|dt| dt.with_timezone(&Utc)) .map_err(|e| { ConversationItemStorageError::StorageError(format!( - "item {fallback_id} invalid created_at: {e}" + "item {fallback_id} invalid {col_created}: {e}" )) })?; @@ -331,8 +366,9 @@ impl ConversationItemStorage for RedisConversationItemStorage { created_at, }; + let si = &self.store.schema.conversation_items; let id_str = conversation_item.id.0.as_str(); - let key = Self::item_key(id_str); + let key = self.item_key(id_str); let mut conn = self .store @@ -343,21 +379,20 @@ impl ConversationItemStorage for RedisConversationItemStorage { let mut pipe = redis::pipe(); - pipe.hset(&key, "id", id_str); + pipe.hset(&key, si.col("id"), id_str); if let Some(rid) = &conversation_item.response_id { - pipe.hset(&key, "response_id", rid); + pipe.hset(&key, si.col("response_id"), rid); } - pipe.hset(&key, "item_type", &conversation_item.item_type); + pipe.hset(&key, si.col("item_type"), &conversation_item.item_type); if let Some(r) = &conversation_item.role { - pipe.hset(&key, "role", r); + pipe.hset(&key, si.col("role"), r); } - pipe.hset(&key, "content", content_json); + pipe.hset(&key, si.col("content"), &content_json); if let Some(s) = &conversation_item.status { - pipe.hset(&key, "status", s); + pipe.hset(&key, si.col("status"), s); } - pipe.hset(&key, "created_at", created_at.to_rfc3339()); + pipe.hset(&key, si.col("created_at"), created_at.to_rfc3339()); - // Expire after configured retention days if let Some(days) = self.store.retention_days { pipe.expire(&key, (days * 24 * 60 * 60) as i64); } @@ -377,7 +412,7 @@ impl ConversationItemStorage for RedisConversationItemStorage { ) -> ConversationItemResult<()> { let cid = conversation_id.0.as_str(); let iid = item_id.0.as_str(); - let key = Self::conv_items_key(cid); + let key = self.conv_items_key(cid); let score = added_at.timestamp_millis() as f64; @@ -400,7 +435,7 @@ impl ConversationItemStorage for RedisConversationItemStorage { params: ListParams, ) -> ConversationItemResult> { let cid = conversation_id.0.as_str(); - let key = Self::conv_items_key(cid); + let key = self.conv_items_key(cid); let mut conn = self .store .pool @@ -410,8 +445,6 @@ impl ConversationItemStorage for RedisConversationItemStorage { let mut min = "-inf".to_string(); let mut max = "+inf".to_string(); - // Track cursor score + id for post-filtering same-millisecond ties, - // matching the composite (added_at, item_id) cursor of Postgres/Oracle. let mut cursor_score: Option = None; let mut cursor_id: Option = None; @@ -423,8 +456,6 @@ impl ConversationItemStorage for RedisConversationItemStorage { if let Some(s) = score { cursor_score = Some(s); cursor_id = Some(after_id.clone()); - // Use inclusive bound so we can post-filter ties by item_id. - // Over-fetch slightly to account for items at the cursor's score. match params.order { SortOrder::Asc => min = s.to_string(), SortOrder::Desc => max = s.to_string(), @@ -432,9 +463,7 @@ impl ConversationItemStorage for RedisConversationItemStorage { } } - // Over-fetch to handle same-score ties that need filtering let fetch_limit = if cursor_score.is_some() { - // Fetch extra to compensate for items we'll filter out at the cursor boundary (params.limit + 32) as isize } else { params.limit as isize @@ -451,10 +480,6 @@ impl ConversationItemStorage for RedisConversationItemStorage { .map_err(|e| ConversationItemStorageError::StorageError(e.to_string()))?, }; - // Post-filter: skip past the cursor item and all same-score predecessors. - // Redis returns same-score members in lexicographic order (ASC) or - // reverse-lex (DESC), so `skip_while` advances past items that appeared - // on the previous page, then `skip(1)` drops the cursor item itself. let item_ids: Vec = if let (Some(_), Some(ref c_id)) = (cursor_score, &cursor_id) { item_ids .into_iter() @@ -470,10 +495,9 @@ impl ConversationItemStorage for RedisConversationItemStorage { return Ok(Vec::::new()); } - // Fetch all items in pipeline let mut pipe = redis::pipe(); for iid in &item_ids { - pipe.hgetall(Self::item_key(iid)); + pipe.hgetall(self.item_key(iid)); } let results: Vec> = pipe @@ -484,11 +508,10 @@ impl ConversationItemStorage for RedisConversationItemStorage { let mut items: Vec = Vec::with_capacity(results.len()); for (i, map) in results.into_iter().enumerate() { if map.is_empty() { - // Item might have been deleted or expired, skip continue; } - items.push(Self::build_item_from_map(&map, &item_ids[i])?); + items.push(self.build_item_from_map(&map, &item_ids[i])?); } Ok(items) @@ -499,7 +522,7 @@ impl ConversationItemStorage for RedisConversationItemStorage { item_id: &ConversationItemId, ) -> ConversationItemResult> { let iid = item_id.0.as_str(); - let key = Self::item_key(iid); + let key = self.item_key(iid); let mut conn = self .store .pool @@ -516,7 +539,7 @@ impl ConversationItemStorage for RedisConversationItemStorage { return Ok(None); } - Self::build_item_from_map(&map, iid).map(Some) + self.build_item_from_map(&map, iid).map(Some) } async fn is_item_linked( @@ -526,7 +549,7 @@ impl ConversationItemStorage for RedisConversationItemStorage { ) -> ConversationItemResult { let cid = conversation_id.0.as_str(); let iid = item_id.0.as_str(); - let key = Self::conv_items_key(cid); + let key = self.conv_items_key(cid); let mut conn = self .store @@ -549,7 +572,7 @@ impl ConversationItemStorage for RedisConversationItemStorage { ) -> ConversationItemResult<()> { let cid = conversation_id.0.as_str(); let iid = item_id.0.as_str(); - let key = Self::conv_items_key(cid); + let key = self.conv_items_key(cid); let mut conn = self .store @@ -574,56 +597,67 @@ impl RedisResponseStorage { Self { store } } - fn response_key(id: &str) -> String { - format!("response:{id}") + fn response_key(&self, id: &str) -> String { + match &self.store.schema.owner { + Some(owner) => format!("{owner}:response:{id}"), + None => format!("response:{id}"), + } } - fn safety_key(identifier: &str) -> String { - format!("safety:{identifier}:responses") + fn safety_key(&self, identifier: &str) -> String { + match &self.store.schema.owner { + Some(owner) => format!("{owner}:safety:{identifier}:responses"), + None => format!("safety:{identifier}:responses"), + } } /// Build a `StoredResponse` from the Redis hash map returned by `HGETALL`. - /// - /// `fallback_id` is used when the map lacks an explicit `"id"` entry - /// (e.g. when the key was derived from the sorted-set member). fn build_response_from_map( + &self, map: std::collections::HashMap, fallback_id: &str, ) -> Result { + let s = &self.store.schema.responses; + + let col_id = s.col("id"); let id = ResponseId( - map.get("id") + map.get(col_id) .cloned() .unwrap_or_else(|| fallback_id.to_string()), ); + let previous_response_id = map - .get("previous_response_id") - .map(|s| ResponseId(s.clone())); - let conversation_id = map.get("conversation_id").cloned(); + .get(s.col("previous_response_id")) + .map(|v| ResponseId(v.clone())); + let conversation_id = map.get(s.col("conversation_id")).cloned(); - let input = parse_json_value(map.get("input").cloned()) + let input = parse_json_value(map.get(s.col("input")).cloned()) .map_err(ResponseStorageError::StorageError)?; - let instructions = map.get("instructions").cloned(); - let output = parse_json_value(map.get("output").cloned()) + let instructions = map.get(s.col("instructions")).cloned(); + let output = parse_json_value(map.get(s.col("output")).cloned()) .map_err(ResponseStorageError::StorageError)?; - let tool_calls = parse_tool_calls(map.get("tool_calls").cloned()) + let tool_calls = parse_tool_calls(map.get(s.col("tool_calls")).cloned()) .map_err(ResponseStorageError::StorageError)?; - let metadata = parse_metadata(map.get("metadata").cloned()) + let metadata = parse_metadata(map.get(s.col("metadata")).cloned()) .map_err(ResponseStorageError::StorageError)?; - let created_at_str = map.get("created_at").ok_or_else(|| { - ResponseStorageError::StorageError(format!("response {fallback_id} missing created_at")) + let col_created = s.col("created_at"); + let created_at_str = map.get(col_created).ok_or_else(|| { + ResponseStorageError::StorageError(format!( + "response {fallback_id} missing {col_created}" + )) })?; let created_at = DateTime::parse_from_rfc3339(created_at_str) .map(|dt| dt.with_timezone(&Utc)) .map_err(|e| { ResponseStorageError::StorageError(format!( - "response {fallback_id} invalid created_at: {e}" + "response {fallback_id} invalid {col_created}: {e}" )) })?; - let safety_identifier = map.get("safety_identifier").cloned(); - let model = map.get("model").cloned(); - let raw_response = parse_raw_response(map.get("raw_response").cloned()) + let safety_identifier = map.get(s.col("safety_identifier")).cloned(); + let model = map.get(s.col("model")).cloned(); + let raw_response = parse_raw_response(map.get(s.col("raw_response")).cloned()) .map_err(ResponseStorageError::StorageError)?; Ok(StoredResponse { @@ -649,9 +683,10 @@ impl ResponseStorage for RedisResponseStorage { &self, response: StoredResponse, ) -> Result { + let sr = &self.store.schema.responses; let response_id = response.id.clone(); let response_id_str = response_id.0.as_str(); - let key = Self::response_key(response_id_str); + let key = self.response_key(response_id_str); let json_input = serde_json::to_string(&response.input)?; let json_output = serde_json::to_string(&response.output)?; @@ -668,30 +703,29 @@ impl ResponseStorage for RedisResponseStorage { let mut pipe = redis::pipe(); - pipe.hset(&key, "id", response_id_str); + pipe.hset(&key, sr.col("id"), response_id_str); if let Some(prev) = &response.previous_response_id { - pipe.hset(&key, "previous_response_id", &prev.0); + pipe.hset(&key, sr.col("previous_response_id"), &prev.0); } - pipe.hset(&key, "input", json_input); + pipe.hset(&key, sr.col("input"), &json_input); if let Some(inst) = &response.instructions { - pipe.hset(&key, "instructions", inst); + pipe.hset(&key, sr.col("instructions"), inst); } - pipe.hset(&key, "output", json_output); - pipe.hset(&key, "tool_calls", json_tool_calls); - pipe.hset(&key, "metadata", json_metadata); - pipe.hset(&key, "created_at", response.created_at.to_rfc3339()); + pipe.hset(&key, sr.col("output"), &json_output); + pipe.hset(&key, sr.col("tool_calls"), &json_tool_calls); + pipe.hset(&key, sr.col("metadata"), &json_metadata); + pipe.hset(&key, sr.col("created_at"), response.created_at.to_rfc3339()); if let Some(safety) = &response.safety_identifier { - pipe.hset(&key, "safety_identifier", safety); + pipe.hset(&key, sr.col("safety_identifier"), safety); } if let Some(model) = &response.model { - pipe.hset(&key, "model", model); + pipe.hset(&key, sr.col("model"), model); } if let Some(cid) = &response.conversation_id { - pipe.hset(&key, "conversation_id", cid); + pipe.hset(&key, sr.col("conversation_id"), cid); } - pipe.hset(&key, "raw_response", json_raw_response); + pipe.hset(&key, sr.col("raw_response"), &json_raw_response); - // Expire after configured retention days if let Some(days) = self.store.retention_days { pipe.expire(&key, (days * 24 * 60 * 60) as i64); } @@ -702,7 +736,7 @@ impl ResponseStorage for RedisResponseStorage { // Index by safety identifier if present if let Some(safety) = &response.safety_identifier { - let safety_key = Self::safety_key(safety); + let safety_key = self.safety_key(safety); let score = response.created_at.timestamp_millis() as f64; conn.zadd::<_, _, _, ()>(safety_key, response_id_str, score) .await @@ -717,7 +751,7 @@ impl ResponseStorage for RedisResponseStorage { response_id: &ResponseId, ) -> Result, ResponseStorageError> { let id = response_id.0.as_str(); - let key = Self::response_key(id); + let key = self.response_key(id); let mut conn = self .store .pool @@ -734,12 +768,15 @@ impl ResponseStorage for RedisResponseStorage { return Ok(None); } - Self::build_response_from_map(map, id).map(Some) + self.build_response_from_map(map, id).map(Some) } async fn delete_response(&self, response_id: &ResponseId) -> ResponseResult<()> { + let sr = &self.store.schema.responses; + let col_safety = sr.col("safety_identifier"); + let id = response_id.0.as_str(); - let key = Self::response_key(id); + let key = self.response_key(id); let mut conn = self .store .pool @@ -747,19 +784,16 @@ impl ResponseStorage for RedisResponseStorage { .await .map_err(|e| ResponseStorageError::StorageError(e.to_string()))?; - // Atomic MULTI/EXEC: read the safety identifier and delete the hash - // in a single transaction so no other client can modify the key - // between the read and the delete. let (safety, ()): (Option, ()) = redis::pipe() .atomic() - .hget(&key, "safety_identifier") + .hget(&key, col_safety) .del(&key) .query_async(&mut conn) .await .map_err(|e| ResponseStorageError::StorageError(e.to_string()))?; if let Some(s) = safety { - conn.zrem::<_, _, ()>(Self::safety_key(&s), id) + conn.zrem::<_, _, ()>(self.safety_key(&s), id) .await .map_err(|e| ResponseStorageError::StorageError(e.to_string()))?; } @@ -767,43 +801,12 @@ impl ResponseStorage for RedisResponseStorage { Ok(()) } - async fn get_response_chain( - &self, - response_id: &ResponseId, - max_depth: Option, - ) -> ResponseResult { - let mut chain = ResponseChain::new(); - let mut current_id = Some(response_id.clone()); - let mut visited = 0usize; - - while let Some(ref lookup_id) = current_id { - if let Some(limit) = max_depth { - if visited >= limit { - break; - } - } - - let fetched = self.get_response(lookup_id).await?; - match fetched { - Some(response) => { - current_id.clone_from(&response.previous_response_id); - chain.responses.push(response); - visited += 1; - } - None => break, - } - } - - chain.responses.reverse(); - Ok(chain) - } - async fn list_identifier_responses( &self, identifier: &str, limit: Option, ) -> ResponseResult> { - let key = Self::safety_key(identifier); + let key = self.safety_key(identifier); let mut conn = self .store .pool @@ -811,7 +814,6 @@ impl ResponseStorage for RedisResponseStorage { .await .map_err(|e| ResponseStorageError::StorageError(e.to_string()))?; - // ZREVRANGE key 0 limit-1 let stop = match limit { Some(l) => (l as isize) - 1, None => -1, @@ -828,7 +830,7 @@ impl ResponseStorage for RedisResponseStorage { let mut pipe = redis::pipe(); for id in &response_ids { - pipe.hgetall(Self::response_key(id)); + pipe.hgetall(self.response_key(id)); } let results: Vec> = pipe @@ -842,14 +844,14 @@ impl ResponseStorage for RedisResponseStorage { continue; } - out.push(Self::build_response_from_map(map, &response_ids[i])?); + out.push(self.build_response_from_map(map, &response_ids[i])?); } Ok(out) } async fn delete_identifier_responses(&self, identifier: &str) -> ResponseResult { - let key = Self::safety_key(identifier); + let key = self.safety_key(identifier); let mut conn = self .store .pool @@ -857,7 +859,6 @@ impl ResponseStorage for RedisResponseStorage { .await .map_err(|e| ResponseStorageError::StorageError(e.to_string()))?; - // Get all IDs let response_ids: Vec = conn .zrange(&key, 0, -1) .await @@ -870,7 +871,7 @@ impl ResponseStorage for RedisResponseStorage { let mut pipe = redis::pipe(); for id in response_ids { - pipe.del(Self::response_key(&id)); + pipe.del(self.response_key(&id)); } pipe.del(&key); diff --git a/data_connector/src/schema.rs b/data_connector/src/schema.rs new file mode 100644 index 0000000000..175dae27e4 --- /dev/null +++ b/data_connector/src/schema.rs @@ -0,0 +1,381 @@ +//! Schema configuration for storage backends. +//! +//! Provides YAML-driven customization of table names and column names. +//! When no schema config is provided, all defaults match the current +//! hardcoded behavior — zero behavioral change. + +use std::collections::HashMap; + +use serde::{Deserialize, Serialize}; + +// ──────────────────────────────────────────────────────────────────────────── +// Types +// ──────────────────────────────────────────────────────────────────────────── + +/// Top-level schema configuration. Drives all SQL generation and key naming. +/// +/// Every field has a default matching current hardcoded behavior, so omitting +/// the entire `schema:` section in YAML produces identical queries to today. +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +#[serde(default)] +pub struct SchemaConfig { + /// Schema owner / key prefix (e.g. `"ADMIN"` for Oracle, `"myapp"` for Redis). + /// The dot in `ADMIN."TABLE"` is generated by `qualified_table()`, not stored here. + #[serde(skip_serializing_if = "Option::is_none")] + pub owner: Option, + + pub conversations: TableConfig, + pub responses: TableConfig, + pub conversation_items: TableConfig, + pub conversation_item_links: TableConfig, +} + +/// Per-table schema configuration. +#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)] +#[serde(default)] +pub struct TableConfig { + /// Physical table name (or Redis key component). + pub table: String, + + /// Column name overrides: `logical_name -> db_column_name`. + /// Fields not listed here use their logical name unchanged. + #[serde(default, skip_serializing_if = "HashMap::is_empty")] + pub columns: HashMap, +} + +// ──────────────────────────────────────────────────────────────────────────── +// Defaults +// ──────────────────────────────────────────────────────────────────────────── + +impl Default for SchemaConfig { + fn default() -> Self { + Self { + owner: None, + conversations: TableConfig::with_table("conversations"), + responses: TableConfig::with_table("responses"), + conversation_items: TableConfig::with_table("conversation_items"), + conversation_item_links: TableConfig::with_table("conversation_item_links"), + } + } +} + +// ──────────────────────────────────────────────────────────────────────────── +// TableConfig methods +// ──────────────────────────────────────────────────────────────────────────── + +impl TableConfig { + /// Create a `TableConfig` with the given table name and no column overrides. + pub fn with_table(name: &str) -> Self { + Self { + table: name.to_string(), + ..Default::default() + } + } + + /// Resolve a column name. + /// + /// Returns the remapped name if an override is configured, or the + /// logical field name unchanged otherwise. + pub fn col<'a>(&'a self, field: &'a str) -> &'a str { + self.columns.get(field).map(String::as_str).unwrap_or(field) + } + + /// Fully qualified table name, e.g. `ADMIN."MY_TABLE"` with an owner + /// or just `my_table` without. + /// + /// The table name is quoted to preserve case. For Oracle, call + /// `SchemaConfig::uppercase_for_oracle()` first so the quoted names match + /// Oracle's uppercase catalog entries. + pub fn qualified_table(&self, owner: Option<&str>) -> String { + match owner { + Some(o) => format!("{o}.\"{}\"", self.table), + None => self.table.clone(), + } + } +} + +// ──────────────────────────────────────────────────────────────────────────── +// Validation +// ──────────────────────────────────────────────────────────────────────────── + +impl SchemaConfig { + /// Uppercase all table names and column override values in-place. + /// + /// Oracle folds unquoted identifiers to uppercase, so existing tables and + /// columns are stored as `CONVERSATIONS`, `CONV_ID`, etc. + /// `qualified_table()` quotes identifiers (`OWNER."table"`), making them + /// case-sensitive. Calling this in `OracleStore::new()` ensures the + /// quoted names match the actual uppercase catalog entries Oracle created. + pub fn uppercase_for_oracle(&mut self) { + for tc in [ + &mut self.conversations, + &mut self.responses, + &mut self.conversation_items, + &mut self.conversation_item_links, + ] { + tc.table.make_ascii_uppercase(); + for val in tc.columns.values_mut() { + val.make_ascii_uppercase(); + } + } + } + + /// Validate the entire schema config at startup. Rejects invalid identifiers. + pub fn validate(&self) -> Result<(), String> { + // Validate owner + if let Some(ref owner) = self.owner { + validate_identifier(owner).map_err(|e| format!("owner: {e}"))?; + } + + // Validate each table config + Self::validate_table("conversations", &self.conversations)?; + Self::validate_table("responses", &self.responses)?; + Self::validate_table("conversation_items", &self.conversation_items)?; + Self::validate_table("conversation_item_links", &self.conversation_item_links)?; + + Ok(()) + } + + fn validate_table(label: &str, tc: &TableConfig) -> Result<(), String> { + validate_identifier(&tc.table).map_err(|e| { + if tc.table.is_empty() { + format!("{label}.table: table name is required (got empty string — did you omit the 'table' key in your config?)") + } else { + format!("{label}.table: {e}") + } + })?; + + for (logical, physical) in &tc.columns { + validate_identifier(logical) + .map_err(|e| format!("{label}.columns key '{logical}': {e}"))?; + validate_identifier(physical) + .map_err(|e| format!("{label}.columns value '{physical}': {e}"))?; + } + + Ok(()) + } +} + +/// Maximum identifier length. Oracle caps at 30 (pre-12.2) or 128, +/// Postgres at 63. We use 128 as a generous upper bound. +const MAX_IDENTIFIER_LEN: usize = 128; + +/// Reject identifiers that are empty, too long, or contain characters outside `[a-zA-Z0-9_]`. +fn validate_identifier(name: &str) -> Result<(), String> { + if name.is_empty() { + return Err("identifier must not be empty".to_string()); + } + if name.len() > MAX_IDENTIFIER_LEN { + return Err(format!( + "identifier '{name}' exceeds maximum length of {MAX_IDENTIFIER_LEN} characters" + )); + } + if !name.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') { + return Err(format!( + "invalid identifier '{name}' — only ASCII alphanumeric and underscores allowed" + )); + } + Ok(()) +} + +// ──────────────────────────────────────────────────────────────────────────── +// Tests +// ──────────────────────────────────────────────────────────────────────────── + +#[cfg(test)] +mod tests { + use super::*; + + // ── Default config ──────────────────────────────────────────────────── + + #[test] + fn default_config_matches_hardcoded_names() { + let cfg = SchemaConfig::default(); + assert_eq!(cfg.conversations.table, "conversations"); + assert_eq!(cfg.responses.table, "responses"); + assert_eq!(cfg.conversation_items.table, "conversation_items"); + assert_eq!(cfg.conversation_item_links.table, "conversation_item_links"); + assert!(cfg.owner.is_none()); + } + + #[test] + fn default_config_validates_successfully() { + SchemaConfig::default() + .validate() + .expect("default config should be valid"); + } + + // ── col() ───────────────────────────────────────────────────────────── + + #[test] + fn col_returns_field_name_when_no_override() { + let tc = TableConfig::with_table("t"); + assert_eq!(tc.col("id"), "id"); + assert_eq!(tc.col("created_at"), "created_at"); + } + + #[test] + fn col_returns_override_when_configured() { + let mut tc = TableConfig::with_table("t"); + tc.columns + .insert("id".to_string(), "CONVERSATION_ID".to_string()); + assert_eq!(tc.col("id"), "CONVERSATION_ID"); + // Non-overridden field still returns itself + assert_eq!(tc.col("created_at"), "created_at"); + } + + // ── qualified_table() ───────────────────────────────────────────────── + + #[test] + fn qualified_table_without_owner() { + let tc = TableConfig::with_table("conversations"); + assert_eq!(tc.qualified_table(None), "conversations"); + } + + #[test] + fn qualified_table_with_owner() { + let tc = TableConfig::with_table("CONVERSATIONS"); + assert_eq!(tc.qualified_table(Some("ADMIN")), "ADMIN.\"CONVERSATIONS\""); + } + + // ── uppercase_for_oracle() ────────────────────────────────────────── + + #[test] + fn uppercase_for_oracle_converts_defaults() { + let mut cfg = SchemaConfig::default(); + cfg.uppercase_for_oracle(); + assert_eq!(cfg.conversations.table, "CONVERSATIONS"); + assert_eq!(cfg.responses.table, "RESPONSES"); + assert_eq!(cfg.conversation_items.table, "CONVERSATION_ITEMS"); + assert_eq!(cfg.conversation_item_links.table, "CONVERSATION_ITEM_LINKS"); + cfg.validate().expect("uppercased config should be valid"); + } + + #[test] + fn uppercase_for_oracle_converts_custom_table_and_columns() { + let mut cfg = SchemaConfig::default(); + cfg.conversations.table = "my_convos".to_string(); + cfg.conversations + .columns + .insert("id".to_string(), "conv_id".to_string()); + cfg.uppercase_for_oracle(); + assert_eq!(cfg.conversations.table, "MY_CONVOS"); + assert_eq!(cfg.conversations.col("id"), "CONV_ID"); + } + + // ── validate() ──────────────────────────────────────────────────────── + + #[test] + fn validate_accepts_valid_identifiers() { + let mut cfg = SchemaConfig { + owner: Some("ADMIN_01".to_string()), + ..Default::default() + }; + cfg.conversations.table = "MY_CONVERSATIONS".to_string(); + cfg.conversations + .columns + .insert("id".to_string(), "CONV_ID".to_string()); + cfg.validate().expect("should be valid"); + } + + #[test] + fn validate_rejects_empty_table_name_with_helpful_message() { + let mut cfg = SchemaConfig::default(); + cfg.conversations.table = String::new(); + let err = cfg.validate().unwrap_err(); + assert!( + err.contains("conversations.table") && err.contains("table name is required"), + "unexpected: {err}" + ); + } + + #[test] + fn validate_rejects_overly_long_identifier() { + let mut cfg = SchemaConfig::default(); + cfg.conversations.table = "a".repeat(129); + let err = cfg.validate().unwrap_err(); + assert!(err.contains("exceeds maximum length"), "unexpected: {err}"); + } + + #[test] + fn validate_rejects_special_characters() { + let mut cfg = SchemaConfig::default(); + cfg.conversations.table = "table;DROP".to_string(); + let err = cfg.validate().unwrap_err(); + assert!( + err.contains("conversations.table") && err.contains("invalid identifier"), + "unexpected: {err}" + ); + } + + #[test] + fn validate_rejects_dots_in_owner() { + let cfg = SchemaConfig { + owner: Some("ADMIN.SCHEMA".to_string()), + ..Default::default() + }; + let err = cfg.validate().unwrap_err(); + assert!( + err.contains("owner") && err.contains("invalid identifier"), + "unexpected: {err}" + ); + } + + #[test] + fn validate_rejects_invalid_column_override_key() { + let mut cfg = SchemaConfig::default(); + cfg.responses + .columns + .insert("bad;col".to_string(), "GOOD_COL".to_string()); + let err = cfg.validate().unwrap_err(); + assert!( + err.contains("columns key") && err.contains("invalid identifier"), + "unexpected: {err}" + ); + } + + #[test] + fn validate_rejects_invalid_column_override_value() { + let mut cfg = SchemaConfig::default(); + cfg.responses + .columns + .insert("id".to_string(), "bad column!".to_string()); + let err = cfg.validate().unwrap_err(); + assert!( + err.contains("columns value") && err.contains("invalid identifier"), + "unexpected: {err}" + ); + } + + // ── Serde roundtrip ─────────────────────────────────────────────────── + + #[test] + fn serde_roundtrip_default() { + let cfg = SchemaConfig::default(); + let json = serde_json::to_string(&cfg).expect("serialize"); + let restored: SchemaConfig = serde_json::from_str(&json).expect("deserialize"); + assert_eq!(cfg, restored); + } + + #[test] + fn serde_roundtrip_custom() { + let mut cfg = SchemaConfig { + owner: Some("ADMIN".to_string()), + ..Default::default() + }; + cfg.conversations.table = "CONVERSATIONS".to_string(); + cfg.conversations + .columns + .insert("id".to_string(), "CONVERSATION_ID".to_string()); + + let json = serde_json::to_string(&cfg).expect("serialize"); + let restored: SchemaConfig = serde_json::from_str(&json).expect("deserialize"); + assert_eq!(cfg, restored); + } + + #[test] + fn serde_deserialize_empty_object_uses_defaults() { + let cfg: SchemaConfig = serde_json::from_str("{}").expect("deserialize empty"); + assert_eq!(cfg, SchemaConfig::default()); + } +} diff --git a/model_gateway/src/main.rs b/model_gateway/src/main.rs index 839ea8d0c6..974734b39a 100644 --- a/model_gateway/src/main.rs +++ b/model_gateway/src/main.rs @@ -854,6 +854,7 @@ impl CliArgs { pool_min, pool_max, pool_timeout_secs, + schema: None, }) } @@ -862,7 +863,11 @@ impl CliArgs { let pool_max = self .postgres_pool_max_size .unwrap_or_else(PostgresConfig::default_pool_max); - let pcf = PostgresConfig { db_url, pool_max }; + let pcf = PostgresConfig { + db_url, + pool_max, + schema: None, + }; pcf.validate().map_err(|e| ConfigError::ValidationFailed { reason: e.to_string(), })?; @@ -883,6 +888,7 @@ impl CliArgs { url, pool_max, retention_days, + schema: None, }; rcf.validate().map_err(|e| ConfigError::ValidationFailed { reason: e.to_string(), diff --git a/model_gateway/tests/routing/test_openai_routing.rs b/model_gateway/tests/routing/test_openai_routing.rs index 074e67fd1a..7093517a33 100644 --- a/model_gateway/tests/routing/test_openai_routing.rs +++ b/model_gateway/tests/routing/test_openai_routing.rs @@ -882,6 +882,7 @@ fn oracle_config_validation_accepts_dsn_only() { pool_min: 1, pool_max: 4, pool_timeout_secs: 30, + schema: None, }) .build_unchecked(); @@ -901,6 +902,7 @@ fn oracle_config_validation_accepts_wallet_alias() { pool_min: 1, pool_max: 8, pool_timeout_secs: 45, + schema: None, }) .build_unchecked(); @@ -922,6 +924,7 @@ fn oracle_config_validation_for_external_auth() { pool_min: 1, pool_max: 4, pool_timeout_secs: 30, + schema: None, }) .build_unchecked(); @@ -940,6 +943,7 @@ fn oracle_config_validation_for_external_auth() { pool_min: 1, pool_max: 4, pool_timeout_secs: 30, + schema: None, }) .build_unchecked();