use crate::error::{Error, Result};
use crate::schema_cache::{Relationship, SchemaCache, Table};
use std::collections::HashMap;
#[derive(Clone, Debug)]
pub struct EmbedPlan {
pub local_column: String,
pub foreign_column: String,
pub foreign_column_type: String,
pub foreign_schema: String,
pub foreign_table: String,
pub is_list: bool,
}
impl EmbedPlan {
pub fn resolve(relationship: &Relationship, schema_cache: &SchemaCache) -> Result<Self> {
let foreign_table_qi = relationship.foreign_table().clone();
let columns = match relationship {
Relationship::ForeignKey { cardinality, .. } => cardinality.columns(),
Relationship::Computed { .. } => {
return Err(Error::EmbeddingError(
"embedding a computed relationship is not supported yet".into(),
))
}
};
if columns.len() != 1 {
return Err(Error::EmbeddingError(format!(
"embedding \"{}\" is not supported yet: it joins on {} columns and \
only single-column joins are implemented",
foreign_table_qi.name,
columns.len()
)));
}
let (local_column, foreign_column) = columns[0].clone();
let foreign_table: &Table = schema_cache.get_table(&foreign_table_qi).ok_or_else(|| {
Error::EmbeddingError(format!(
"cannot embed \"{}\": it is not in an exposed schema",
foreign_table_qi
))
})?;
let foreign_column_type = foreign_table
.get_column(&foreign_column)
.map(|c| c.nominal_type.clone())
.ok_or_else(|| {
Error::EmbeddingError(format!(
"cannot embed \"{}\": join column \"{}\" not found",
foreign_table_qi, foreign_column
))
})?;
Ok(Self {
local_column,
foreign_column,
foreign_column_type,
foreign_schema: foreign_table_qi.schema.clone(),
foreign_table: foreign_table_qi.name.clone(),
is_list: !relationship.is_to_one(),
})
}
pub fn children_sql(&self, limit: Option<i64>, columns: &[String]) -> Result<String> {
let type_name = castable_type_name(&self.foreign_column_type).ok_or_else(|| {
Error::EmbeddingError(format!(
"cannot embed \"{}\": join column type \"{}\" is not a plain type name",
self.foreign_table, self.foreign_column_type
))
})?;
let projection = self.projection(columns);
let mut inner = format!(
"SELECT {} FROM {}.{} WHERE {} = ANY($1::{}[])",
projection,
postrust_sql::escape_ident(&self.foreign_schema),
postrust_sql::escape_ident(&self.foreign_table),
postrust_sql::escape_ident(&self.foreign_column),
type_name
);
if let Some(limit) = limit {
inner.push_str(&format!(" LIMIT {}", limit));
}
Ok(format!("SELECT row_to_json(t) FROM ({}) t", inner))
}
pub fn children_grouped_sql(&self, limit: Option<i64>, columns: &[String]) -> Result<String> {
let type_name = castable_type_name(&self.foreign_column_type).ok_or_else(|| {
Error::EmbeddingError(format!(
"cannot embed \"{}\": join column type \"{}\" is not a plain type name",
self.foreign_table, self.foreign_column_type
))
})?;
let key = postrust_sql::escape_ident(&self.foreign_column);
let mut inner = format!(
"SELECT {} FROM {}.{} WHERE {} = ANY($1::{}[])",
self.projection(columns),
postrust_sql::escape_ident(&self.foreign_schema),
postrust_sql::escape_ident(&self.foreign_table),
key,
type_name
);
if let Some(limit) = limit {
inner.push_str(&format!(" LIMIT {}", limit));
}
Ok(format!(
"SELECT to_jsonb(c.{key}) AS k, json_agg(row_to_json(c)) AS v \
FROM ({inner}) c GROUP BY c.{key}",
key = key,
inner = inner
))
}
pub fn embed_expression(
&self,
parent_alias: &str,
child_alias: &str,
inner_select: &str,
limit: Option<i64>,
) -> Result<String> {
let mut inner = format!(
"SELECT {} FROM {}.{} AS {} WHERE {}.{} = {}.{}",
inner_select,
postrust_sql::escape_ident(&self.foreign_schema),
postrust_sql::escape_ident(&self.foreign_table),
postrust_sql::escape_ident(child_alias),
postrust_sql::escape_ident(child_alias),
postrust_sql::escape_ident(&self.foreign_column),
postrust_sql::escape_ident(parent_alias),
postrust_sql::escape_ident(&self.local_column),
);
if let Some(limit) = limit {
inner.push_str(&format!(" LIMIT {}", limit));
} else if !self.is_list {
inner.push_str(" LIMIT 1");
}
let alias = postrust_sql::escape_ident(&format!("{}_j", child_alias));
Ok(if self.is_list {
format!(
"COALESCE((SELECT json_agg(row_to_json({alias})) FROM ({inner}) {alias}), '[]'::json)",
alias = alias,
inner = inner
)
} else {
format!(
"(SELECT row_to_json({alias}) FROM ({inner}) {alias})",
alias = alias,
inner = inner
)
})
}
fn projection(&self, columns: &[String]) -> String {
if columns.is_empty() {
return "*".to_string();
}
let mut wanted: Vec<&str> = Vec::with_capacity(columns.len() + 1);
for column in columns {
if !wanted.contains(&column.as_str()) {
wanted.push(column);
}
}
if !wanted.contains(&self.foreign_column.as_str()) {
wanted.push(&self.foreign_column);
}
wanted
.into_iter()
.map(postrust_sql::escape_ident)
.collect::<Vec<_>>()
.join(", ")
}
}
fn castable_type_name(pg_type: &str) -> Option<&str> {
if pg_type.is_empty() {
return None;
}
if pg_type
.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '_')
{
Some(pg_type)
} else {
None
}
}
pub fn key_to_text(value: &serde_json::Value) -> Option<String> {
match value {
serde_json::Value::Null => None,
serde_json::Value::String(s) => Some(s.clone()),
other => Some(other.to_string()),
}
}
pub fn group_from_aggregated(
rows: Vec<(serde_json::Value, serde_json::Value)>,
) -> HashMap<String, Vec<serde_json::Value>> {
let mut grouped: HashMap<String, Vec<serde_json::Value>> = HashMap::with_capacity(rows.len());
for (key, children) in rows {
let key = key_to_text(&key).unwrap_or_default();
let children = match children {
serde_json::Value::Array(items) => items,
_ => Vec::new(),
};
grouped.entry(key).or_default().extend(children);
}
grouped
}
pub fn group_by_key(
children: Vec<serde_json::Value>,
foreign_column: &str,
) -> HashMap<String, Vec<serde_json::Value>> {
let mut grouped: HashMap<String, Vec<serde_json::Value>> = HashMap::new();
for child in children {
let key = child
.get(foreign_column)
.and_then(key_to_text)
.unwrap_or_default();
grouped.entry(key).or_default().push(child);
}
grouped
}
pub fn attach_to_parent(
parent: &mut serde_json::Value,
field_name: &str,
plan: &EmbedPlan,
grouped: &HashMap<String, Vec<serde_json::Value>>,
) {
let key = parent.get(&plan.local_column).and_then(key_to_text);
let matches = key
.as_ref()
.and_then(|k| grouped.get(k))
.cloned()
.unwrap_or_default();
let value = if plan.is_list {
serde_json::Value::Array(matches)
} else {
matches
.into_iter()
.next()
.unwrap_or(serde_json::Value::Null)
};
if let Some(object) = parent.as_object_mut() {
object.insert(field_name.to_string(), value);
}
}
pub fn parent_keys(parents: &[serde_json::Value], local_column: &str) -> Vec<String> {
let mut seen = std::collections::HashSet::new();
let mut keys = Vec::new();
for parent in parents {
if let Some(key) = parent.get(local_column).and_then(key_to_text) {
if seen.insert(key.clone()) {
keys.push(key);
}
}
}
keys
}
#[cfg(test)]
mod tests {
use super::*;
fn plan(is_list: bool) -> EmbedPlan {
EmbedPlan {
local_column: "id".into(),
foreign_column: "user_id".into(),
foreign_column_type: "int4".into(),
foreign_schema: "public".into(),
foreign_table: "posts".into(),
is_list,
}
}
#[test]
fn children_sql_binds_keys_as_a_cast_array() {
let sql = plan(true).children_sql(None, &[]).unwrap();
assert_eq!(
sql,
"SELECT row_to_json(t) FROM (SELECT * FROM \"public\".\"posts\" \
WHERE \"user_id\" = ANY($1::int4[])) t"
);
}
#[test]
fn embed_expression_aggregates_a_to_many_relation() {
let sql = plan(true)
.embed_expression("p", "posts", r#""id", "title""#, None)
.unwrap();
assert!(sql.starts_with("COALESCE((SELECT json_agg("), "{}", sql);
assert!(sql.contains(r#""posts"."user_id" = "p"."id""#), "{}", sql);
assert!(
sql.contains(r#"AS "posts""#),
"the child table is aliased: {}",
sql
);
assert!(!sql.contains("ANY("), "{}", sql);
assert!(sql.contains("'[]'::json"), "{}", sql);
}
#[test]
fn embed_expression_takes_one_row_for_a_to_one_relation() {
let sql = plan(false)
.embed_expression("p", "author", r#""id""#, None)
.unwrap();
assert!(sql.contains("row_to_json"), "{}", sql);
assert!(!sql.contains("json_agg"), "{}", sql);
assert!(
sql.contains("LIMIT 1"),
"a to-one relation yields one row: {}",
sql
);
}
#[test]
fn embed_expression_limits_rows_per_parent() {
let sql = plan(true)
.embed_expression("p", "posts", r#""id""#, Some(25))
.unwrap();
assert!(sql.contains("LIMIT 25"), "{}", sql);
}
#[test]
fn children_sql_projects_only_the_requested_columns() {
let sql = plan(true)
.children_sql(None, &["title".to_string(), "body".to_string()])
.unwrap();
assert!(
sql.contains(r#"SELECT "title", "body", "user_id" FROM"#),
"{}",
sql
);
assert!(
!sql.contains("SELECT *"),
"an unrequested column should not be read at all: {}",
sql
);
}
#[test]
fn children_sql_always_includes_the_join_column() {
let sql = plan(true)
.children_sql(None, &["title".to_string()])
.unwrap();
assert!(sql.contains(r#""user_id""#), "{}", sql);
}
#[test]
fn children_sql_does_not_repeat_the_join_column() {
let sql = plan(true)
.children_sql(None, &["user_id".to_string(), "title".to_string()])
.unwrap();
assert_eq!(
sql.matches(r#""user_id""#).count(),
2,
"expected the column once in the projection and once in the WHERE: {}",
sql
);
}
#[test]
fn children_sql_escapes_column_names() {
let sql = plan(true)
.children_sql(None, &[r#"ev"il"#.to_string()])
.unwrap();
assert!(sql.contains(r#""ev""il""#), "{}", sql);
}
#[test]
fn children_sql_falls_back_to_every_column() {
let sql = plan(true).children_sql(None, &[]).unwrap();
assert!(sql.contains("SELECT * FROM"), "{}", sql);
}
#[test]
fn children_sql_applies_a_limit() {
let sql = plan(true).children_sql(Some(25), &[]).unwrap();
assert!(sql.contains("LIMIT 25"), "{}", sql);
}
#[test]
fn children_sql_rejects_a_non_plain_type_name() {
let mut p = plan(true);
p.foreign_column_type = "int4; DROP TABLE users".into();
assert!(p.children_sql(None, &[]).is_err());
}
#[test]
fn keys_are_rendered_without_json_quoting() {
assert_eq!(key_to_text(&serde_json::json!(7)), Some("7".to_string()));
assert_eq!(
key_to_text(&serde_json::json!("abc")),
Some("abc".to_string())
);
assert_eq!(key_to_text(&serde_json::Value::Null), None);
}
#[test]
fn parent_keys_are_distinct_and_skip_nulls() {
let parents = vec![
serde_json::json!({"id": 1}),
serde_json::json!({"id": 2}),
serde_json::json!({"id": 1}),
serde_json::json!({"id": null}),
];
assert_eq!(parent_keys(&parents, "id"), vec!["1", "2"]);
}
#[test]
fn to_many_attaches_an_array_and_empty_when_absent() {
let grouped = group_by_key(vec![serde_json::json!({"id": 10, "user_id": 1})], "user_id");
let mut matched = serde_json::json!({"id": 1});
attach_to_parent(&mut matched, "posts", &plan(true), &grouped);
assert_eq!(matched["posts"].as_array().map(|a| a.len()), Some(1));
let mut unmatched = serde_json::json!({"id": 2});
attach_to_parent(&mut unmatched, "posts", &plan(true), &grouped);
assert_eq!(
unmatched["posts"],
serde_json::json!([]),
"a to-many with no matches must still be an array"
);
}
#[test]
fn to_one_attaches_an_object_or_null() {
let grouped = group_by_key(vec![serde_json::json!({"id": 10, "user_id": 1})], "user_id");
let mut matched = serde_json::json!({"id": 1});
attach_to_parent(&mut matched, "author", &plan(false), &grouped);
assert_eq!(matched["author"]["id"], serde_json::json!(10));
let mut unmatched = serde_json::json!({"id": 2});
attach_to_parent(&mut unmatched, "author", &plan(false), &grouped);
assert_eq!(unmatched["author"], serde_json::Value::Null);
}
}