use serde::Deserialize;
use serde_json::{json, Value};
use trusty_common::memory_core::store::kg::ExpandDirection;
use crate::service::core_kg::{DEFAULT_KG_LIST_LIMIT, MAX_KG_LIST_LIMIT};
use crate::transport::api_error::ApiError;
use crate::AppState;
use super::{to_value, NoParams, PalaceParams};
pub use crate::service::{DreamStatusPayload, KgGraphPayload, KgNeighborsPayload, KgSeedPayload};
fn default_kg_list_limit() -> usize {
DEFAULT_KG_LIST_LIMIT
}
#[derive(Debug, Deserialize)]
pub struct KgListParams {
pub palace_id: String,
#[serde(default = "default_kg_list_limit")]
pub limit: usize,
#[serde(default)]
pub offset: usize,
}
pub async fn kg_subjects_with_counts(
state: &AppState,
params: KgListParams,
) -> Result<Value, ApiError> {
let limit = params.limit.clamp(1, MAX_KG_LIST_LIMIT);
let rows = crate::service::MemoryService::new(state.clone())
.kg_list_subjects_with_counts(¶ms.palace_id, limit)
.await?;
Ok(Value::Array(
rows.into_iter()
.map(|(subject, count)| json!({ "subject": subject, "count": count }))
.collect(),
))
}
pub async fn kg_all(state: &AppState, params: KgListParams) -> Result<Value, ApiError> {
let limit = params.limit.clamp(1, MAX_KG_LIST_LIMIT);
to_value(
crate::service::MemoryService::new(state.clone())
.kg_list_all(¶ms.palace_id, limit, params.offset)
.await?,
)
}
pub async fn kg_count(state: &AppState, params: PalaceParams) -> Result<Value, ApiError> {
let active = crate::service::MemoryService::new(state.clone())
.kg_count(¶ms.palace_id)
.await?;
Ok(json!({ "active": active }))
}
pub async fn kg_graph(state: &AppState, params: PalaceParams) -> Result<Value, ApiError> {
to_value(
crate::service::MemoryService::new(state.clone())
.kg_graph(¶ms.palace_id)
.await?,
)
}
const DEFAULT_KG_SEED_LIMIT: usize = 75;
const MAX_KG_SEED_LIMIT: usize = 200;
fn default_kg_seed_limit() -> usize {
DEFAULT_KG_SEED_LIMIT
}
#[derive(Debug, Deserialize)]
pub struct KgSeedParams {
pub palace_id: String,
#[serde(default = "default_kg_seed_limit")]
pub limit: usize,
}
pub async fn kg_graph_seed(state: &AppState, params: KgSeedParams) -> Result<Value, ApiError> {
let limit = params.limit.clamp(1, MAX_KG_SEED_LIMIT);
to_value(
crate::service::MemoryService::new(state.clone())
.kg_graph_seed(¶ms.palace_id, limit)
.await?,
)
}
const MAX_KG_NEIGHBOR_HOPS: usize = 4;
fn default_kg_neighbor_hops() -> usize {
1
}
#[derive(Debug, Deserialize)]
pub struct KgNeighborsParams {
pub palace_id: String,
pub node: String,
#[serde(default)]
pub direction: Option<String>,
#[serde(default = "default_kg_neighbor_hops")]
pub max_hops: usize,
}
pub async fn kg_graph_neighbors(
state: &AppState,
params: KgNeighborsParams,
) -> Result<Value, ApiError> {
let direction = match params.direction.as_deref().unwrap_or("both") {
"in" | "inbound" => ExpandDirection::In,
"out" | "outbound" => ExpandDirection::Out,
"both" => ExpandDirection::Both,
other => {
return Err(ApiError::bad_request(format!(
"direction must be in|out|both, got {other:?}"
)))
}
};
let max_hops = params.max_hops.clamp(1, MAX_KG_NEIGHBOR_HOPS);
to_value(
crate::service::MemoryService::new(state.clone())
.kg_neighbors(¶ms.palace_id, ¶ms.node, direction, max_hops)
.await?,
)
}
const TRIPLE_ID_SEPARATOR: u8 = 0x00;
pub fn encode_triple_id(subject: &str, predicate: &str, object: &str) -> Result<String, String> {
use base64::Engine as _;
for (field, value) in [
("subject", subject),
("predicate", predicate),
("object", object),
] {
if value.as_bytes().contains(&TRIPLE_ID_SEPARATOR) {
return Err(format!(
"{field} must not contain the null-byte separator (\\0); got {value:?}"
));
}
}
let mut buf = Vec::with_capacity(subject.len() + predicate.len() + object.len() + 2);
buf.extend_from_slice(subject.as_bytes());
buf.push(TRIPLE_ID_SEPARATOR);
buf.extend_from_slice(predicate.as_bytes());
buf.push(TRIPLE_ID_SEPARATOR);
buf.extend_from_slice(object.as_bytes());
Ok(base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(&buf))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TripleIdError {
Malformed,
LegacyPair,
}
pub fn decode_triple_id(id: &str) -> Result<(String, String, String), TripleIdError> {
use base64::Engine as _;
let bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(id)
.map_err(|_| TripleIdError::Malformed)?;
let parts = bytes
.split(|&b| b == TRIPLE_ID_SEPARATOR)
.map(|part| String::from_utf8(part.to_vec()))
.collect::<Result<Vec<_>, _>>()
.map_err(|_| TripleIdError::Malformed)?;
match <[String; 3]>::try_from(parts) {
Ok([subject, predicate, object]) => Ok((subject, predicate, object)),
Err(parts) if parts.len() == 2 => Err(TripleIdError::LegacyPair),
Err(_) => Err(TripleIdError::Malformed),
}
}
#[derive(Debug, Deserialize)]
pub struct DeleteTripleParams {
pub palace_id: String,
pub triple_id: String,
}
pub async fn kg_delete_triple(
state: &AppState,
params: DeleteTripleParams,
) -> Result<Value, ApiError> {
let (subject, predicate, object) = match decode_triple_id(¶ms.triple_id) {
Ok(triple) => triple,
Err(TripleIdError::LegacyPair) => {
return Err(ApiError::bad_request(
"triple id names only (subject, predicate) — it must encode \
base64url(subject\\0predicate\\0object) so the delete targets one triple",
))
}
Err(TripleIdError::Malformed) => {
return Err(ApiError::not_found(
"invalid triple id — expected base64url(subject\\0predicate\\0object)",
))
}
};
let closed = crate::service::MemoryService::new(state.clone())
.kg_retract_triple(¶ms.palace_id, &subject, &predicate, &object)
.await?;
if closed > 0 {
Ok(json!({ "closed": closed }))
} else {
Err(ApiError::not_found(format!(
"no active triple with subject={subject:?} predicate={predicate:?} \
object={object:?} in palace {:?}",
params.palace_id
)))
}
}
pub async fn dream_status(state: &AppState, _params: NoParams) -> Result<Value, ApiError> {
to_value(
crate::service::MemoryService::new(state.clone())
.dream_status_aggregate()
.await,
)
}
pub async fn palace_dream_status(
state: &AppState,
params: PalaceParams,
) -> Result<Value, ApiError> {
to_value(
crate::service::MemoryService::new(state.clone())
.dream_status_for_palace(¶ms.palace_id)
.await?,
)
}
pub async fn dream_run(state: &AppState, _params: NoParams) -> Result<Value, ApiError> {
to_value(
crate::service::MemoryService::new(state.clone())
.dream_run()
.await?,
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn decode_triple_id_round_trips() {
let id = encode_triple_id("s", "p", "o").expect("encode");
assert_eq!(
decode_triple_id(&id).expect("decode"),
("s".into(), "p".into(), "o".into())
);
}
#[test]
fn encode_triple_id_rejects_null_byte() {
assert!(encode_triple_id("a\0b", "p", "o").is_err());
}
#[test]
fn decode_triple_id_rejects_the_legacy_pair_form() {
use base64::Engine as _;
let legacy = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(b"s\0p");
assert_eq!(
decode_triple_id(&legacy).expect_err("a pair is not a triple"),
TripleIdError::LegacyPair
);
}
#[test]
fn decode_triple_id_returns_none_for_invalid_input() {
assert_eq!(
decode_triple_id("!!!not base64!!!").expect_err("not decodable"),
TripleIdError::Malformed
);
}
}