use super::system::{is_physical_shard_collection, is_protected_collection, AppState};
use crate::{
error::DbError,
server::response::ApiResponse,
storage::{http_client::get_http_client, query_cache},
sync::{LogEntry, Operation},
transaction::TransactionId,
triggers::{fire_collection_triggers, TriggerEvent},
};
use axum::{
extract::{Path, Query, State},
http::{HeaderMap, StatusCode},
response::Json,
};
use serde::Deserialize;
use serde_json::Value;
pub fn get_transaction_id(headers: &HeaderMap) -> Option<TransactionId> {
headers
.get("X-Transaction-ID")
.and_then(|h| h.to_str().ok())
.and_then(|s| {
let id_str = s.strip_prefix("tx:").unwrap_or(s);
id_str.parse::<u64>().ok()
})
.map(TransactionId::from_u64)
}
async fn inject_auto_embeddings_if_needed(
storage: &std::sync::Arc<crate::storage::StorageEngine>,
db_name: &str,
collection: &crate::storage::Collection,
mut data: serde_json::Value,
) -> Result<serde_json::Value, DbError> {
use crate::server::llm_client::LLMClient;
use crate::storage::index::VectorIndexConfig;
let configs: Vec<VectorIndexConfig> = collection.get_all_vector_index_configs();
if configs.is_empty() {
return Ok(data);
}
if !data.is_object() {
return Ok(data);
}
let obj = data.as_object_mut().unwrap();
for config in configs {
if let Some(ref source_field) = config.embedding_source {
let target_field = config.field.clone();
if let Some(existing) = obj.get(&target_field) {
if let Some(arr) = existing.as_array() {
if arr.len() == config.dimension {
continue;
}
}
}
let text = match obj.get(source_field).and_then(|v| v.as_str()) {
Some(t) if !t.trim().is_empty() => t.to_string(),
_ => continue,
};
let provider = config
.embedding_provider
.clone()
.unwrap_or_else(|| "openai".to_string());
let client = match LLMClient::from_storage(
storage,
db_name,
Some(&provider),
config.embedding_model.clone(),
) {
Ok(c) => c,
Err(e) => {
tracing::warn!(
"Auto-embedding skipped for vector index '{}' (provider {}): {}",
config.name,
provider,
e
);
continue;
}
};
let emb = match client.embed(&text).await {
Ok(e) => e,
Err(e) => {
tracing::warn!(
"Auto-embedding failed for field '{}' in vector index '{}': {}",
source_field,
config.name,
e
);
continue;
}
};
if emb.len() != config.dimension {
tracing::warn!(
"Auto-embedding dimension mismatch for index '{}': provider returned {}, \
expected {} — storing document without vector",
config.name,
emb.len(),
config.dimension
);
continue;
}
obj.insert(target_field, serde_json::json!(emb));
}
}
Ok(data)
}
#[derive(Debug, Deserialize)]
pub struct CopyShardRequest {
pub source_address: String,
}
pub async fn insert_document(
State(state): State<AppState>,
Path((db_name, coll_name)): Path<(String, String)>,
headers: HeaderMap,
Json(data): Json<Value>,
) -> Result<Json<Value>, DbError> {
let database = state.storage.get_database(&db_name)?;
let collection = match database.get_collection(&coll_name) {
Ok(coll) => coll,
Err(DbError::CollectionNotFound(_)) => {
tracing::info!(
"Auto-creating document collection {}/{}",
db_name,
coll_name
);
database.create_collection(coll_name.clone(), None)?;
database.get_collection(&coll_name)?
}
Err(e) => return Err(e),
};
let data =
inject_auto_embeddings_if_needed(&state.storage, &db_name, &collection, data).await?;
if let Some(tx_id) = get_transaction_id(&headers) {
let tx_manager = state.storage.transaction_manager()?;
let tx_arc = tx_manager.get(tx_id)?;
let mut tx = tx_arc
.write()
.map_err(|_| DbError::InternalError("Transaction lock poisoned".into()))?;
let wal = tx_manager.wal().clone();
let lock_manager = tx_manager.lock_manager().clone();
let doc = collection.insert_tx(&mut tx, &wal, &lock_manager, data)?;
return Ok(Json(doc.to_value()));
}
if let Some(shard_config) = collection.get_shard_config() {
tracing::info!(
"[INSERT] shard_config found: num_shards={}",
shard_config.num_shards
);
if shard_config.num_shards > 0 {
if let Some(ref coordinator) = state.shard_coordinator {
if !headers.contains_key("X-Shard-Direct") {
tracing::info!(
"[INSERT] Using ShardCoordinator for {}/{}",
db_name,
coll_name
);
let doc = coordinator
.insert(&db_name, &coll_name, &shard_config, data)
.await?;
return Ok(Json(doc));
}
} else {
tracing::error!(
"[INSERT] Sharded collection {}/{} but no shard_coordinator available!",
db_name,
coll_name
);
return Err(DbError::InternalError(
"Sharded collection requires ShardCoordinator".to_string(),
));
}
}
}
let doc = collection.insert(data)?;
let is_shard = is_physical_shard_collection(&coll_name);
if !is_shard {
if let Some(ref log) = state.replication_log {
let entry = LogEntry {
sequence: 0,
node_id: "".to_string(),
database: db_name.clone(),
collection: coll_name.clone(),
operation: Operation::Insert,
key: doc.key.clone(),
data: serde_json::to_vec(&doc.to_value()).ok(),
timestamp: chrono::Utc::now().timestamp_millis() as u64,
origin_sequence: None,
};
let _ = log.append(entry);
}
}
query_cache::get_query_cache().invalidate_collection(&coll_name);
if !coll_name.starts_with('_') {
let notifier = state.queue_worker.as_ref().map(|w| w.notifier());
let _ = fire_collection_triggers(
&state.storage,
notifier.as_ref(),
&db_name,
&coll_name,
TriggerEvent::Insert,
&doc,
None,
);
}
Ok(Json(doc.to_value()))
}
pub async fn insert_documents_batch(
State(state): State<AppState>,
Path((db_name, coll_name)): Path<(String, String)>,
headers: HeaderMap,
Json(documents): Json<Vec<Value>>,
) -> Result<Json<Value>, DbError> {
let database = state.storage.get_database(&db_name)?;
let collection = database.get_collection(&coll_name)?;
if !headers.contains_key("X-Shard-Direct") {
return Err(DbError::BadRequest(
"Batch endpoint requires X-Shard-Direct header".to_string(),
));
}
let is_physical_shard = coll_name.contains("_s")
&& coll_name
.chars()
.last()
.map(|c| c.is_ascii_digit())
.unwrap_or(false);
let insert_count = if is_physical_shard {
let keyed_docs: Vec<(String, Value)> = documents
.iter()
.map(|doc| {
let key = doc
.get("_key")
.and_then(|k| k.as_str())
.unwrap_or("")
.to_string();
(key, doc.clone())
})
.filter(|(key, _)| !key.is_empty())
.collect();
collection.upsert_batch(keyed_docs)?
} else {
collection.insert_batch(documents.clone())?.len()
};
query_cache::get_query_cache().invalidate_collection(&coll_name);
let is_migration = headers.contains_key("X-Migration");
if is_migration {
tracing::debug!(
"BATCH: Skipping replica forwarding - migration operation for {}/{}",
db_name,
coll_name
);
} else if let Some(ref coordinator) = state.shard_coordinator {
let is_rebalancing = coordinator.is_rebalancing();
if is_rebalancing {
tracing::debug!(
"BATCH: Skipping replica forwarding during rebalancing for {}/{}",
db_name,
coll_name
);
} else {
if let Some(idx) = coll_name.rfind("_s") {
let base_coll = &coll_name[..idx];
if let Ok(shard_id) = coll_name[idx + 2..].parse::<u16>() {
if let Some(table) = coordinator.get_shard_table(&db_name, base_coll) {
if let Some(assignment) = table.assignments.get(&shard_id) {
if !assignment.replica_nodes.is_empty() {
let client = get_http_client();
let secret = state.cluster_secret();
if let Some(ref cluster_manager) = state.cluster_manager {
let mut futures = Vec::new();
for replica_node in &assignment.replica_nodes {
if let Some(addr) =
cluster_manager.get_node_api_address(replica_node)
{
let url = format!(
"http://{}/_api/database/{}/document/{}/_replica",
addr, db_name, coll_name
);
tracing::debug!("REPLICA FWD: Forwarding {} docs to replica {} at {}", documents.len(), replica_node, addr);
let client = client.clone();
let secret = secret.clone();
let docs = documents.clone();
let future = async move {
let _ = tokio::time::timeout(
std::time::Duration::from_secs(10), client
.post(&url)
.header("X-Shard-Direct", "true")
.header("X-Cluster-Secret", &secret)
.json(&docs)
.send(),
)
.await;
};
futures.push(future);
}
}
tokio::spawn(async move {
futures::future::join_all(futures).await;
});
}
}
}
}
}
}
}
}
Ok(Json(serde_json::json!({
"inserted": insert_count,
"success": true
})))
}
pub async fn insert_documents_replica(
State(state): State<AppState>,
Path((db_name, coll_name)): Path<(String, String)>,
headers: HeaderMap,
Json(documents): Json<Vec<Value>>,
) -> Result<Json<Value>, DbError> {
if !headers.contains_key("X-Shard-Direct") {
return Err(DbError::BadRequest(
"Replica endpoint requires X-Shard-Direct header".to_string(),
));
}
let database = state.storage.get_database(&db_name)?;
let collection = database.get_collection(&coll_name)?;
let keyed_docs: Vec<(String, Value)> = documents
.iter()
.map(|doc| {
let key = doc
.get("_key")
.and_then(|k| k.as_str())
.unwrap_or("")
.to_string();
(key, doc.clone())
})
.filter(|(key, _)| !key.is_empty())
.collect();
let insert_count = collection.upsert_batch(keyed_docs)?;
query_cache::get_query_cache().invalidate_collection(&coll_name);
tracing::debug!(
"REPLICA: Stored {} docs for {}/{}",
insert_count,
db_name,
coll_name
);
Ok(Json(serde_json::json!({
"inserted": insert_count,
"success": true
})))
}
pub async fn verify_documents_exist(
State(state): State<AppState>,
Path((db_name, coll_name)): Path<(String, String)>,
Json(request): Json<serde_json::Value>,
) -> Result<Json<Value>, DbError> {
let keys = request
.get("keys")
.and_then(|k| k.as_array())
.ok_or_else(|| DbError::BadRequest("Missing 'keys' array in request body".to_string()))?;
let database = state.storage.get_database(&db_name)?;
let collection = database.get_collection(&coll_name)?;
let mut found: Vec<String> = Vec::new();
let mut missing: Vec<String> = Vec::new();
for key_value in keys {
if let Some(key) = key_value.as_str() {
match collection.get(key) {
Ok(_) => found.push(key.to_string()),
Err(_) => missing.push(key.to_string()),
}
}
}
let total_checked = found.len() + missing.len();
tracing::debug!(
"VERIFY: Checked {} docs in {}/{}: {} found, {} missing",
total_checked,
db_name,
coll_name,
found.len(),
missing.len()
);
Ok(Json(serde_json::json!({
"found": found,
"missing": missing,
"total_checked": total_checked
})))
}
pub async fn copy_shard_data(
State(state): State<AppState>,
Path((db_name, coll_name)): Path<(String, String)>,
Json(request): Json<CopyShardRequest>,
) -> Result<Json<Value>, DbError> {
tracing::info!(
"COPY_SHARD: Copying {}/{} from {}",
db_name,
coll_name,
request.source_address
);
let secret = state.cluster_secret();
let client = get_http_client();
let meta_url = format!(
"http://{}/_api/database/{}/collection/{}",
request.source_address, db_name, coll_name
);
let meta_res = client
.get(&meta_url)
.header("X-Cluster-Secret", &secret)
.header("X-Shard-Direct", "true")
.timeout(std::time::Duration::from_secs(10))
.send()
.await;
let mut source_count = 0;
let mut check_count = false;
if let Ok(res) = meta_res {
if res.status().is_success() {
if let Ok(json) = res.json::<serde_json::Value>().await {
if let Some(c) = json.get("count").and_then(|v| v.as_u64()) {
source_count = c as usize;
check_count = true;
}
}
}
}
let database = state.storage.get_database(&db_name)?;
let collection = match database.get_collection(&coll_name) {
Ok(c) => c,
Err(_) => {
database.create_collection(coll_name.clone(), None)?;
database.get_collection(&coll_name)?
}
};
if check_count {
let local_count = collection.count();
if local_count == source_count {
tracing::info!(
"COPY_SHARD: Skipping sync for {}/{} (Count match: {})",
db_name,
coll_name,
local_count
);
return Ok(Json(serde_json::json!({
"copied": 0,
"success": true,
"skipped": true
})));
}
tracing::info!(
"COPY_SHARD: Count mismatch for {}/{} (Local: {}, Source: {}). Truncating before sync.",
db_name,
coll_name,
local_count,
source_count
);
let _ = collection.truncate();
}
let url = format!(
"http://{}/_api/database/{}/cursor",
request.source_address, db_name
);
let query = format!("FOR doc IN {} RETURN doc", coll_name);
let res = client
.post(&url)
.header("X-Cluster-Secret", &secret)
.json(&serde_json::json!({ "query": query }))
.timeout(std::time::Duration::from_secs(120))
.send()
.await
.map_err(|e| DbError::InternalError(format!("Request failed: {}", e)))?;
if !res.status().is_success() {
let status = res.status();
let body_text = res
.text()
.await
.unwrap_or_else(|_| "Could not read error body".to_string());
tracing::error!(
"COPY_SHARD: Source query failed. Status: {}, Body: {}",
status,
body_text
);
return Err(DbError::InternalError(format!(
"Source query failed: {}. Body: {}",
status, body_text
)));
}
let body: serde_json::Value = res
.json()
.await
.map_err(|e| DbError::InternalError(format!("Parse failed: {}", e)))?;
let docs = body
.get("result")
.and_then(|r| r.as_array())
.ok_or_else(|| DbError::InternalError("No result array".to_string()))?;
let keyed_docs: Vec<(String, serde_json::Value)> = docs
.iter()
.map(|doc| {
let key = doc
.get("_key")
.and_then(|k| k.as_str())
.unwrap_or("")
.to_string();
(key, doc.clone())
})
.filter(|(key, _)| !key.is_empty())
.collect();
let count = keyed_docs.len();
collection.upsert_batch(keyed_docs)?;
tracing::info!(
"COPY_SHARD: Copied {} docs to {}/{}",
count,
db_name,
coll_name
);
Ok(Json(serde_json::json!({
"copied": count,
"success": true
})))
}
pub async fn get_document(
State(state): State<AppState>,
Path((db_name, coll_name, key)): Path<(String, String, String)>,
headers: HeaderMap,
) -> Result<ApiResponse<Value>, DbError> {
let database = state.storage.get_database(&db_name)?;
let collection = database.get_collection(&coll_name)?;
if let Some(shard_config) = collection.get_shard_config() {
if shard_config.num_shards > 0 {
if let Some(ref coordinator) = state.shard_coordinator {
let doc = coordinator.get(&db_name, &coll_name, &key).await?;
let mut doc_value = doc;
let replicas = coordinator.get_replicas(&key, &shard_config);
if let Value::Object(ref mut map) = doc_value {
map.insert("_replicas".to_string(), serde_json::json!(replicas));
}
return Ok(ApiResponse::new(doc_value, &headers));
}
}
}
let doc = collection.get(&key)?;
Ok(ApiResponse::new(doc.to_value(), &headers))
}
pub async fn update_document(
State(state): State<AppState>,
Path((db_name, coll_name, key)): Path<(String, String, String)>,
headers: HeaderMap,
Query(params): Query<std::collections::HashMap<String, String>>,
Json(mut data): Json<Value>,
) -> Result<Json<Value>, DbError> {
let database = state.storage.get_database(&db_name)?;
let collection = database.get_collection(&coll_name)?;
let upsert = params.get("upsert").map(|v| v == "true").unwrap_or(false);
if let Some(tx_id) = get_transaction_id(&headers) {
let tx_manager = state.storage.transaction_manager()?;
let tx_arc = tx_manager.get(tx_id)?;
let mut tx = tx_arc
.write()
.map_err(|_| DbError::InternalError("Transaction lock poisoned".into()))?;
let wal = tx_manager.wal().clone();
let lock_manager = tx_manager.lock_manager().clone();
let doc = collection.update_tx(&mut tx, &wal, &lock_manager, &key, data)?;
return Ok(Json(doc.to_value()));
}
if let Some(shard_config) = collection.get_shard_config() {
if shard_config.num_shards > 0 {
if let Some(ref coordinator) = state.shard_coordinator {
if !headers.contains_key("X-Shard-Direct") {
let doc = coordinator
.update(&db_name, &coll_name, &shard_config, &key, data)
.await?;
return Ok(Json(doc));
}
}
}
}
let old_doc_value = collection.get(&key).ok().map(|d| d.to_value());
let if_match = headers
.get(axum::http::header::IF_MATCH)
.and_then(|v| v.to_str().ok())
.map(|s| s.trim_matches('"').to_string());
let (doc, was_upsert) = match if_match {
Some(rev) => (collection.update_with_rev(&key, &rev, data.clone())?, false),
None => match collection.update(&key, data.clone()) {
Ok(doc) => (doc, false),
Err(DbError::DocumentNotFound(_)) if upsert => {
if let Value::Object(ref mut obj) = data {
obj.insert("_key".to_string(), Value::String(key.clone()));
}
(collection.insert(data)?, true)
}
Err(e) => return Err(e),
},
};
let is_shard = is_physical_shard_collection(&coll_name);
let is_sharded_logical = collection.get_shard_config().is_some();
if !is_shard && !is_sharded_logical {
if let Some(ref log) = state.replication_log {
let entry = LogEntry {
sequence: 0,
node_id: "".to_string(),
database: db_name.clone(),
collection: coll_name.clone(),
operation: Operation::Update,
key: doc.key.clone(),
data: serde_json::to_vec(&doc.to_value()).ok(),
timestamp: chrono::Utc::now().timestamp_millis() as u64,
origin_sequence: None,
};
let _ = log.append(entry);
}
}
query_cache::get_query_cache().invalidate_collection(&coll_name);
if !coll_name.starts_with('_') {
let notifier = state.queue_worker.as_ref().map(|w| w.notifier());
let event = if was_upsert {
TriggerEvent::Insert
} else {
TriggerEvent::Update
};
let _ = fire_collection_triggers(
&state.storage,
notifier.as_ref(),
&db_name,
&coll_name,
event,
&doc,
old_doc_value.as_ref(),
);
}
Ok(Json(doc.to_value()))
}
pub async fn delete_document(
State(state): State<AppState>,
Path((db_name, coll_name, key)): Path<(String, String, String)>,
headers: HeaderMap,
) -> Result<StatusCode, DbError> {
if is_protected_collection(&db_name, &coll_name) {
return Err(DbError::BadRequest(format!(
"Cannot delete documents from protected collection: {}",
coll_name
)));
}
let database = state.storage.get_database(&db_name)?;
let collection = database.get_collection(&coll_name)?;
if let Some(tx_id) = get_transaction_id(&headers) {
let tx_manager = state.storage.transaction_manager()?;
let tx_arc = tx_manager.get(tx_id)?;
let mut tx = tx_arc
.write()
.map_err(|_| DbError::InternalError("Transaction lock poisoned".into()))?;
let wal = tx_manager.wal().clone();
let lock_manager = tx_manager.lock_manager().clone();
collection.delete_tx(&mut tx, &wal, &lock_manager, &key)?;
return Ok(StatusCode::NO_CONTENT);
}
if let Some(shard_config) = collection.get_shard_config() {
if shard_config.num_shards > 0 {
if let Some(ref coordinator) = state.shard_coordinator {
if !headers.contains_key("X-Shard-Direct") {
coordinator
.delete(&db_name, &coll_name, &shard_config, &key)
.await?;
return Ok(StatusCode::NO_CONTENT);
}
}
}
}
let old_doc = collection.get(&key).ok();
collection.delete(&key)?;
query_cache::get_query_cache().invalidate_collection(&coll_name);
if collection.get_type() == "blob" {
tracing::info!(
"Compacting blob collection {}/{} after deletion of {}",
db_name,
coll_name,
key
);
collection.compact();
}
let is_shard = is_physical_shard_collection(&coll_name);
let is_sharded_logical = collection.get_shard_config().is_some();
if !is_shard && !is_sharded_logical {
if let Some(ref log) = state.replication_log {
let entry = LogEntry {
sequence: 0,
node_id: state
.cluster_manager
.as_ref()
.map(|m| m.local_node_id())
.unwrap_or_else(|| "".to_string()),
database: db_name.clone(),
collection: coll_name.clone(),
operation: Operation::Delete,
key: key.clone(),
data: None,
timestamp: chrono::Utc::now().timestamp_millis() as u64,
origin_sequence: None,
};
let _ = log.append(entry);
}
}
if !coll_name.starts_with('_') {
if let Some(old_doc) = old_doc {
let notifier = state.queue_worker.as_ref().map(|w| w.notifier());
let old_doc_value = old_doc.to_value();
let _ = fire_collection_triggers(
&state.storage,
notifier.as_ref(),
&db_name,
&coll_name,
TriggerEvent::Delete,
&old_doc,
Some(&old_doc_value),
);
}
}
Ok(StatusCode::NO_CONTENT)
}