use std::collections::BTreeMap;
use crate::broker::RequestContext;
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct AppliedContext {
pub tenant_id: String,
pub project_id: String,
pub purpose: String,
pub scopes: String,
pub correlation_id: String,
pub user_id: String,
pub service_identity: String,
pub decision_id: String,
pub attributes: BTreeMap<String, String>,
}
impl AppliedContext {
pub fn from_request(ctx: &RequestContext) -> Self {
Self {
tenant_id: ctx.tenant_id.clone(),
project_id: ctx.project_id.clone(),
purpose: ctx.purpose.clone(),
scopes: ctx.scopes.join(","),
correlation_id: ctx.correlation_id.clone(),
user_id: ctx.user_id.clone(),
service_identity: ctx.service_identity.clone(),
decision_id: ctx.decision_id.clone(),
attributes: BTreeMap::new(),
}
}
pub fn is_empty(&self) -> bool {
self.tenant_id.is_empty()
&& self.project_id.is_empty()
&& self.purpose.is_empty()
&& self.scopes.is_empty()
&& self.correlation_id.is_empty()
&& self.user_id.is_empty()
&& self.service_identity.is_empty()
&& self.decision_id.is_empty()
&& self.attributes.is_empty()
}
pub fn session_context_pairs(&self) -> Vec<(&'static str, &str)> {
vec![
("app.current_tenant_id", self.tenant_id.as_str()),
("app.current_project_id", self.project_id.as_str()),
("app.current_purpose", self.purpose.as_str()),
("app.current_scopes", self.scopes.as_str()),
("app.current_correlation_id", self.correlation_id.as_str()),
("app.current_user_id", self.user_id.as_str()),
(
"app.current_service_identity",
self.service_identity.as_str(),
),
("app.current_decision_id", self.decision_id.as_str()),
]
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ContextEffect {
Enforced { mechanism: String },
Advisory { recorded_in: String },
Unsupported { reason: String },
}
impl ContextEffect {
pub fn is_enforced(&self) -> bool {
matches!(self, Self::Enforced { .. })
}
pub fn is_unsupported(&self) -> bool {
matches!(self, Self::Unsupported { .. })
}
}
pub fn enforce_with_mechanism(ctx: &AppliedContext, mechanism: &str) -> ContextEffect {
if ctx.is_empty() {
ContextEffect::Advisory {
recorded_in: "no_context_to_apply".into(),
}
} else {
ContextEffect::Enforced {
mechanism: mechanism.into(),
}
}
}
pub fn escape_sql_string(value: &str, _dialect: SqlDialect) -> String {
value.replace('\'', "''")
}
pub trait BackendContextEnforcer: Send + Sync {
fn backend_label(&self) -> &str;
fn enforce(&self, ctx: &AppliedContext) -> ContextEffect;
}
pub fn render_sql_session_settings(ctx: &AppliedContext, dialect: SqlDialect) -> Vec<String> {
if ctx.is_empty() {
return Vec::new();
}
let mut out = Vec::new();
let pairs = ctx.session_context_pairs();
for (key, value) in pairs {
if value.is_empty() {
continue;
}
match dialect {
SqlDialect::Postgres => {
let escaped = escape_sql_string(value, dialect);
out.push(format!("SET LOCAL {key} = '{escaped}'"));
}
SqlDialect::Mysql => {
let var_name = key.replace('.', "_");
let escaped = escape_sql_string(value, dialect);
out.push(format!("SET @{var_name} = '{escaped}'"));
}
SqlDialect::Sqlite => {
let escaped_key = escape_sql_string(key, dialect);
let escaped_value = escape_sql_string(value, dialect);
out.push(format!(
"INSERT OR REPLACE INTO _udb_context(key, value) \
VALUES ('{escaped_key}', '{escaped_value}')"
));
}
SqlDialect::Clickhouse => {
let escaped_key = key.replace('.', "_");
let escaped_value = escape_sql_string(value, dialect);
out.push(format!("SET {escaped_key} = '{escaped_value}'"));
}
SqlDialect::Mssql => {
let key_id = key.replace('.', "_");
let escaped_key = escape_sql_string(&key_id, dialect);
let escaped_value = escape_sql_string(value, dialect);
out.push(format!(
"EXEC sp_set_session_context @key = N'{escaped_key}', \
@value = N'{escaped_value}'"
));
}
}
}
out
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SqlDialect {
Postgres,
Mysql,
Sqlite,
Clickhouse,
Mssql,
}
pub fn render_document_filter_prefix(ctx: &AppliedContext) -> Option<serde_json::Value> {
if ctx.tenant_id.is_empty() && ctx.project_id.is_empty() {
return None;
}
let mut prefix = serde_json::Map::new();
if !ctx.tenant_id.is_empty() {
prefix.insert(
"_tenant_id".into(),
serde_json::Value::String(ctx.tenant_id.clone()),
);
}
if !ctx.project_id.is_empty() {
prefix.insert(
"_project_id".into(),
serde_json::Value::String(ctx.project_id.clone()),
);
}
Some(serde_json::Value::Object(prefix))
}
pub fn render_kv_key_prefix(ctx: &AppliedContext) -> String {
if ctx.tenant_id.is_empty() && ctx.project_id.is_empty() {
return String::new();
}
let tenant = if ctx.tenant_id.is_empty() {
"default"
} else {
ctx.tenant_id.as_str()
};
let project = if ctx.project_id.is_empty() {
"default"
} else {
ctx.project_id.as_str()
};
format!("t:{tenant}/p:{project}/")
}
pub fn render_cypher_context_parameters(
ctx: &AppliedContext,
) -> std::collections::HashMap<String, serde_json::Value> {
let mut params = std::collections::HashMap::new();
if !ctx.tenant_id.is_empty() {
params.insert(
"ctx_tenant_id".to_string(),
serde_json::Value::String(ctx.tenant_id.clone()),
);
}
if !ctx.project_id.is_empty() {
params.insert(
"ctx_project_id".to_string(),
serde_json::Value::String(ctx.project_id.clone()),
);
}
if !ctx.purpose.is_empty() {
params.insert(
"ctx_purpose".to_string(),
serde_json::Value::String(ctx.purpose.clone()),
);
}
params
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn empty_context_renders_no_sql() {
let ctx = AppliedContext::default();
assert!(render_sql_session_settings(&ctx, SqlDialect::Postgres).is_empty());
assert!(render_sql_session_settings(&ctx, SqlDialect::Mysql).is_empty());
}
#[test]
fn postgres_uses_set_local_with_double_single_escape() {
let ctx = AppliedContext {
tenant_id: "acme'corp".into(),
project_id: "p1".into(),
..Default::default()
};
let stmts = render_sql_session_settings(&ctx, SqlDialect::Postgres);
assert_eq!(stmts.len(), 2);
assert!(stmts[0].contains("SET LOCAL app.current_tenant_id = 'acme''corp'"));
assert!(stmts[1].contains("SET LOCAL app.current_project_id = 'p1'"));
}
#[test]
fn postgres_session_settings_are_transaction_scoped_for_pool_safety() {
let ctx = AppliedContext {
tenant_id: "acme".into(),
project_id: "p1".into(),
..Default::default()
};
let stmts = render_sql_session_settings(&ctx, SqlDialect::Postgres);
assert!(
!stmts.is_empty(),
"a populated context must render session settings"
);
for stmt in &stmts {
let upper = stmt.to_ascii_uppercase();
assert!(
stmt.starts_with("SET LOCAL "),
"PG session setting must be transaction-scoped (SET LOCAL) for \
transaction-pool safety, got: {stmt}"
);
assert!(
!upper.contains("SET SESSION"),
"PG session setting must never be session-scoped (would leak \
across a transaction-pooled connection): {stmt}"
);
}
}
#[test]
fn mysql_munges_dots_to_underscores_and_uses_user_vars() {
let ctx = AppliedContext {
tenant_id: "acme".into(),
project_id: "p1".into(),
..Default::default()
};
let stmts = render_sql_session_settings(&ctx, SqlDialect::Mysql);
assert!(stmts[0].contains("SET @app_current_tenant_id = 'acme'"));
assert!(stmts[1].contains("SET @app_current_project_id = 'p1'"));
}
#[test]
fn sqlite_inserts_into_context_temp_table() {
let ctx = AppliedContext {
tenant_id: "acme".into(),
project_id: "p1".into(),
..Default::default()
};
let stmts = render_sql_session_settings(&ctx, SqlDialect::Sqlite);
assert!(stmts[0].contains("INSERT OR REPLACE INTO _udb_context"));
assert!(stmts[0].contains("VALUES ('app.current_tenant_id', 'acme')"));
}
#[test]
fn mssql_emits_sp_set_session_context_with_nvarchar_literals() {
let ctx = AppliedContext {
tenant_id: "acme'corp".into(), project_id: "p1".into(),
..Default::default()
};
let stmts = render_sql_session_settings(&ctx, SqlDialect::Mssql);
assert_eq!(stmts.len(), 2);
assert!(
stmts[0].contains(
"EXEC sp_set_session_context @key = N'app_current_tenant_id', \
@value = N'acme''corp'"
),
"got: {}",
stmts[0]
);
assert!(stmts[1].contains("@key = N'app_current_project_id'"));
assert!(stmts[1].contains("@value = N'p1'"));
assert!(!stmts.iter().any(|stmt| stmt.contains("@read_only")));
}
#[test]
fn document_filter_prefix_emits_underscore_fields() {
let ctx = AppliedContext {
tenant_id: "acme".into(),
project_id: "p1".into(),
..Default::default()
};
let prefix = render_document_filter_prefix(&ctx).unwrap();
assert_eq!(prefix["_tenant_id"], "acme");
assert_eq!(prefix["_project_id"], "p1");
}
#[test]
fn kv_key_prefix_uses_default_when_field_missing() {
let ctx = AppliedContext {
tenant_id: "acme".into(),
..Default::default()
};
let prefix = render_kv_key_prefix(&ctx);
assert_eq!(prefix, "t:acme/p:default/");
}
#[test]
fn kv_key_prefix_empty_when_nothing_set() {
let ctx = AppliedContext::default();
assert_eq!(render_kv_key_prefix(&ctx), "");
}
#[test]
fn cypher_parameters_skip_empty_fields() {
let ctx = AppliedContext {
tenant_id: "acme".into(),
..Default::default()
};
let params = render_cypher_context_parameters(&ctx);
assert_eq!(params.len(), 1);
assert_eq!(params["ctx_tenant_id"], "acme");
assert!(!params.contains_key("ctx_project_id"));
}
#[test]
fn context_effect_classifies_correctly() {
let enforced = ContextEffect::Enforced {
mechanism: "SET LOCAL".into(),
};
assert!(enforced.is_enforced());
assert!(!enforced.is_unsupported());
let advisory = ContextEffect::Advisory {
recorded_in: "tx_metadata".into(),
};
assert!(!advisory.is_enforced());
assert!(!advisory.is_unsupported());
let unsupported = ContextEffect::Unsupported {
reason: "backend has no session settings".into(),
};
assert!(!unsupported.is_enforced());
assert!(unsupported.is_unsupported());
}
}