use std::time::{Duration, Instant};
use tonic::{Request, Response, Status};
use uuid::Uuid;
use crate::proto::udb::core::embedding::services::v1 as embedding_pb;
use crate::proto::{VectorHybridSearchRequest, VectorSearchRequest, VectorSet};
use crate::runtime::channels::OperationChannel;
use super::super::native_helpers::{
admit_on as native_admit_on, native_next_page_token, native_offset_page_window,
native_service_context, non_empty_json, validate_request_tenant,
};
use super::EmbeddingServiceImpl;
use super::config::{
EMBEDDING_SOURCE_MSG, MAX_SOURCES_PER_TENANT, STATUS_ACTIVE, STATUS_DELETED,
TOPIC_BACKFILL_REQUESTED, TOPIC_SOURCE_DELETED, TOPIC_SOURCE_REGISTERED, resolve_top_k,
retrieve_fusion_weights, retrieve_score_threshold,
};
use super::errors::{
embedding_field_violation, embedding_required_field, embedding_source_not_found_status,
require_source_tenant_column, validate_register_source_required_fields,
validate_report_embedding_required_fields, validate_reported_vector,
};
use super::model::{build_embedding_point, merge_retrieve_filter, stored_source_from_json};
use super::store::{active_sources_read, source_conflict, source_read_by_name, source_record};
impl EmbeddingServiceImpl {
pub(crate) fn resolve_source_tenant_column(
&self,
project_id: &str,
source_message_type: &str,
) -> Result<(String, String), Status> {
let catalog = self.require_catalog()?;
let state = catalog.active_for(project_id);
let manifest = &state.manifest;
let table = crate::broker::table_for_message(manifest, source_message_type).ok_or_else(
|| {
embedding_field_violation(
"source_message_type",
"must be present in the active catalog manifest",
format!(
"source_message_type '{source_message_type}' is not present in the active catalog \
manifest"
),
)
},
)?;
let resolved = crate::runtime::postgres_helpers::tenant_column_ref(table)
.map(|column| column.column_name.clone());
let tenant_column = require_source_tenant_column(resolved, source_message_type)?;
let cdc_topic = table.cdc_topic.clone();
Ok((tenant_column, cdc_topic))
}
pub(crate) async fn upsert_reported_embedding(
&self,
project_id: &str,
collection: &str,
tenant_id: &str,
source_name: &str,
row_pk: &str,
vector: Vec<f32>,
dims: i32,
) -> Result<(), Status> {
validate_reported_vector(dims, &vector)?;
let runtime = self.require_runtime()?;
let dim = if dims > 0 { dims } else { vector.len() as i32 };
let point = build_embedding_point(row_pk, vector, tenant_id, source_name)?;
let vector_instance = runtime
.choose_instance_name_for_project("qdrant", true, project_id)
.map(str::to_string)
.unwrap_or_else(|| "default".to_string());
runtime
.vector_upsert_backend_target(
Some(&vector_instance),
project_id,
collection,
dim,
vec![point],
)
.await
}
}
pub(crate) fn retrieve_hit_payload_json(payload: Option<&prost_types::Struct>) -> String {
let Some(payload) = payload else {
return String::new();
};
let mut json = crate::runtime::executor_utils::struct_to_json(payload);
let Some(object) = json.as_object_mut() else {
return String::new();
};
object.remove("_tenant_id");
object.remove("_source");
if let Some(parent) = object.remove("_parent_pk") {
object.insert("parent_pk".to_string(), parent);
}
if let Some(seq) = object.remove("_chunk_seq") {
object.insert("chunk_seq".to_string(), seq);
}
if object.is_empty() {
return String::new();
}
serde_json::Value::Object(std::mem::take(object)).to_string()
}
pub(crate) fn parse_grpc_timeout(value: &str) -> Option<Duration> {
let value = value.trim();
if value.len() < 2 {
return None;
}
let (digits, unit) = value.split_at(value.len() - 1);
let amount: u64 = digits.parse().ok()?;
let nanos = match unit {
"H" => amount.checked_mul(3_600_000_000_000)?,
"M" => amount.checked_mul(60_000_000_000)?,
"S" => amount.checked_mul(1_000_000_000)?,
"m" => amount.checked_mul(1_000_000)?,
"u" => amount.checked_mul(1_000)?,
"n" => amount,
_ => return None,
};
Some(Duration::from_nanos(nanos))
}
pub(crate) fn remaining_before_deadline(
deadline: Option<Instant>,
now: Instant,
) -> Result<Option<Duration>, Status> {
match deadline {
None => Ok(None),
Some(deadline) => match deadline.checked_duration_since(now) {
Some(remaining) if !remaining.is_zero() => Ok(Some(remaining)),
_ => Err(crate::runtime::executor_utils::deadline_exceeded_status(
"embedding",
"retrieve",
crate::runtime::executor_utils::HTTP_RETRYABLE_BACKOFF_MS,
"retrieve deadline exceeded before semantic search dispatch",
)),
},
}
}
pub(crate) async fn register_source(
svc: &EmbeddingServiceImpl,
request: Request<embedding_pb::RegisterSourceRequest>,
) -> Result<Response<embedding_pb::RegisterSourceResponse>, Status> {
let metadata = request.metadata().clone();
let req = request.into_inner();
validate_request_tenant(&metadata, &req.tenant_id)?;
let tenant_id = req.tenant_id.trim().to_string();
let source_name = req.source_name.trim().to_string();
let source_message_type = req.source_message_type.trim().to_string();
let target_collection = req.target_collection.trim().to_string();
validate_register_source_required_fields(&source_name, &source_message_type)?;
if target_collection.is_empty() {
return Err(embedding_required_field(
"target_collection",
"must be a non-empty target vector collection",
"target_collection is required",
));
}
let _admit = native_admit_on(
svc.channels.as_ref(),
&svc.metrics,
"embedding",
OperationChannel::Admin,
&tenant_id,
None,
)
.await?;
let runtime = svc.require_runtime()?;
let context = native_service_context(&metadata, &tenant_id, "");
let (tenant_column, source_cdc_topic) =
svc.resolve_source_tenant_column(&context.project_id, &source_message_type)?;
let existing = runtime
.native_entity_read_for_service(
"embedding",
&context,
source_read_by_name(&tenant_id, &source_name),
)
.await?
.first()
.map(stored_source_from_json);
if existing.is_none() {
let active = runtime
.native_entity_read_for_service(
"embedding",
&context,
active_sources_read(&tenant_id, 0, (MAX_SOURCES_PER_TENANT as u32) + 1),
)
.await?;
if active.len() >= MAX_SOURCES_PER_TENANT {
return Err(crate::runtime::executor_utils::quota_refusal_status(
"embedding",
"tenant embedding-source quota",
format!("tenant embedding-source quota exhausted ({MAX_SOURCES_PER_TENANT})"),
));
}
}
let source_id = existing
.as_ref()
.map(|row| row.source_id.clone())
.filter(|id| !id.trim().is_empty())
.unwrap_or_else(|| Uuid::new_v4().to_string());
let text_fields: Vec<String> = req
.text_fields
.iter()
.map(|field| field.trim().to_string())
.filter(|field| !field.is_empty())
.collect();
let text_fields_json = serde_json::to_string(&text_fields).unwrap_or_else(|_| "[]".to_string());
let model_id = req.model_id.trim().to_string();
let metadata_json = non_empty_json(&req.metadata_json);
runtime
.native_entity_write_for_service(
"embedding",
&context,
EMBEDDING_SOURCE_MSG,
source_record(
&source_id,
&tenant_id,
&source_name,
&source_message_type,
&text_fields_json,
&target_collection,
&model_id,
&tenant_column,
&source_cdc_topic,
STATUS_ACTIVE,
&metadata_json,
),
source_conflict(),
)
.await?;
svc.emit_source_event(
TOPIC_SOURCE_REGISTERED,
&tenant_id,
&context.project_id,
&source_name,
serde_json::json!({
"source_message_type": source_message_type,
"target_collection": target_collection,
"model_id": model_id,
"tenant_column": tenant_column,
}),
)
.await;
if let Some(prev) = existing.as_ref() {
let model_changed = prev.model_id.trim() != model_id;
let collection_changed = prev.target_collection.trim() != target_collection;
if model_changed || collection_changed {
let reindex_id = Uuid::new_v4().to_string();
svc.emit_source_event(
TOPIC_BACKFILL_REQUESTED,
&tenant_id,
&context.project_id,
&source_name,
serde_json::json!({
"backfill_id": reindex_id,
"source_message_type": source_message_type,
"target_collection": target_collection,
"model_id": model_id,
"reason": "model_or_collection_changed",
"previous_model_id": prev.model_id,
"previous_collection": prev.target_collection,
}),
)
.await;
}
}
Ok(Response::new(embedding_pb::RegisterSourceResponse {
source_id,
source_name,
tenant_column,
message: "embedding source registered".to_string(),
error: None,
}))
}
pub(crate) async fn list_sources(
svc: &EmbeddingServiceImpl,
request: Request<embedding_pb::ListSourcesRequest>,
) -> Result<Response<embedding_pb::ListSourcesResponse>, Status> {
let metadata = request.metadata().clone();
let req = request.into_inner();
validate_request_tenant(&metadata, &req.tenant_id)?;
let tenant_id = req.tenant_id.trim().to_string();
let _admit = native_admit_on(
svc.channels.as_ref(),
&svc.metrics,
"embedding",
OperationChannel::Read,
&tenant_id,
None,
)
.await?;
let runtime = svc.require_runtime()?;
let context = native_service_context(&metadata, &tenant_id, "");
let page_window = native_offset_page_window(
1,
req.page_size,
&req.page_token,
MAX_SOURCES_PER_TENANT as i32,
);
let rows = runtime
.native_entity_read_for_service(
"embedding",
&context,
active_sources_read(
&tenant_id,
page_window.offset as u64,
(page_window.limit as u32).min(MAX_SOURCES_PER_TENANT as u32),
),
)
.await?;
let sources = rows
.iter()
.map(stored_source_from_json)
.map(|source| embedding_pb::EmbeddingSourceSummary {
source_id: source.source_id,
source_name: source.source_name,
source_message_type: source.source_message_type,
target_collection: source.target_collection,
model_id: source.model_id,
status: source.status,
})
.collect::<Vec<_>>();
let next_page_token =
native_next_page_token(page_window.offset, page_window.limit, sources.len());
Ok(Response::new(embedding_pb::ListSourcesResponse {
sources,
message: "ok".to_string(),
error: None,
next_page_token,
}))
}
pub(crate) async fn delete_source(
svc: &EmbeddingServiceImpl,
request: Request<embedding_pb::DeleteSourceRequest>,
) -> Result<Response<embedding_pb::DeleteSourceResponse>, Status> {
let metadata = request.metadata().clone();
let req = request.into_inner();
validate_request_tenant(&metadata, &req.tenant_id)?;
let tenant_id = req.tenant_id.trim().to_string();
let source_name = req.source_name.trim().to_string();
if source_name.is_empty() {
return Err(embedding_required_field(
"source_name",
"must be a non-empty embedding source name",
"source_name is required",
));
}
let _admit = native_admit_on(
svc.channels.as_ref(),
&svc.metrics,
"embedding",
OperationChannel::Admin,
&tenant_id,
None,
)
.await?;
let runtime = svc.require_runtime()?;
let context = native_service_context(&metadata, &tenant_id, "");
let stored = runtime
.native_entity_read_for_service(
"embedding",
&context,
source_read_by_name(&tenant_id, &source_name),
)
.await?
.first()
.map(stored_source_from_json);
let Some(stored) = stored.filter(|row| row.status != STATUS_DELETED) else {
return Ok(Response::new(embedding_pb::DeleteSourceResponse {
deleted: true,
message: "embedding source not found".to_string(),
error: None,
}));
};
runtime
.native_entity_write_for_service(
"embedding",
&context,
EMBEDDING_SOURCE_MSG,
source_record(
&stored.source_id,
&tenant_id,
&stored.source_name,
&stored.source_message_type,
&stored.text_fields_json,
&stored.target_collection,
&stored.model_id,
&stored.tenant_column,
&stored.source_cdc_topic,
STATUS_DELETED,
"{}",
),
source_conflict(),
)
.await?;
svc.emit_source_event(
TOPIC_SOURCE_DELETED,
&tenant_id,
&context.project_id,
&source_name,
serde_json::json!({ "target_collection": stored.target_collection }),
)
.await;
Ok(Response::new(embedding_pb::DeleteSourceResponse {
deleted: true,
message: "embedding source deleted".to_string(),
error: None,
}))
}
pub(crate) async fn backfill(
svc: &EmbeddingServiceImpl,
request: Request<embedding_pb::BackfillRequest>,
) -> Result<Response<embedding_pb::BackfillResponse>, Status> {
let metadata = request.metadata().clone();
let req = request.into_inner();
validate_request_tenant(&metadata, &req.tenant_id)?;
let tenant_id = req.tenant_id.trim().to_string();
let source_name = req.source_name.trim().to_string();
if source_name.is_empty() {
return Err(embedding_required_field(
"source_name",
"must be a non-empty embedding source name",
"source_name is required",
));
}
let _admit = native_admit_on(
svc.channels.as_ref(),
&svc.metrics,
"embedding",
OperationChannel::Admin,
&tenant_id,
None,
)
.await?;
let runtime = svc.require_runtime()?;
let context = native_service_context(&metadata, &tenant_id, "");
let stored = runtime
.native_entity_read_for_service(
"embedding",
&context,
source_read_by_name(&tenant_id, &source_name),
)
.await?
.first()
.map(stored_source_from_json)
.filter(|source| source.status == STATUS_ACTIVE)
.ok_or_else(|| embedding_source_not_found_status("backfill"))?;
let backfill_id = Uuid::new_v4().to_string();
svc.emit_source_event(
TOPIC_BACKFILL_REQUESTED,
&tenant_id,
&context.project_id,
&source_name,
serde_json::json!({
"backfill_id": backfill_id,
"source_message_type": stored.source_message_type,
"target_collection": stored.target_collection,
"model_id": stored.model_id,
}),
)
.await;
Ok(Response::new(embedding_pb::BackfillResponse {
backfill_id,
accepted: true,
message: "embedding backfill requested".to_string(),
error: None,
}))
}
pub(crate) async fn report_embedding(
svc: &EmbeddingServiceImpl,
request: Request<embedding_pb::ReportEmbeddingRequest>,
) -> Result<Response<embedding_pb::ReportEmbeddingResponse>, Status> {
let metadata = request.metadata().clone();
let req = request.into_inner();
validate_request_tenant(&metadata, &req.tenant_id)?;
let tenant_id = req.tenant_id.trim().to_string();
let source_name = req.source_name.trim().to_string();
let row_pk = req.row_pk.trim().to_string();
validate_report_embedding_required_fields(&source_name, &row_pk)?;
validate_reported_vector(req.dims, &req.vector)?;
let _admit = native_admit_on(
svc.channels.as_ref(),
&svc.metrics,
"embedding",
OperationChannel::Vector,
&tenant_id,
None,
)
.await?;
let runtime = svc.require_runtime()?;
let context = native_service_context(&metadata, &tenant_id, "");
let stored = runtime
.native_entity_read_for_service(
"embedding",
&context,
source_read_by_name(&tenant_id, &source_name),
)
.await?
.first()
.map(stored_source_from_json)
.filter(|source| source.status == STATUS_ACTIVE)
.ok_or_else(|| embedding_source_not_found_status("report_embedding"))?;
svc.upsert_reported_embedding(
&context.project_id,
&stored.collection(),
&tenant_id,
&source_name,
&row_pk,
req.vector,
req.dims,
)
.await?;
Ok(Response::new(embedding_pb::ReportEmbeddingResponse {
upserted: true,
message: "embedding upserted".to_string(),
error: None,
}))
}
pub(crate) async fn retrieve(
svc: &EmbeddingServiceImpl,
request: Request<embedding_pb::RetrieveRequest>,
) -> Result<Response<embedding_pb::RetrieveResponse>, Status> {
let metadata = request.metadata().clone();
let deadline = metadata
.get("grpc-timeout")
.and_then(|value| value.to_str().ok())
.and_then(parse_grpc_timeout)
.map(|budget| Instant::now() + budget);
let req = request.into_inner();
validate_request_tenant(&metadata, &req.tenant_id)?;
let tenant_id = req.tenant_id.trim().to_string();
let source_name = req.source_name.trim().to_string();
if source_name.is_empty() {
return Err(embedding_required_field(
"source_name",
"must be a non-empty embedding source name",
"source_name is required",
));
}
if req.query_vector.is_empty() {
return Err(embedding_required_field(
"query_vector",
"must contain at least one embedding dimension",
"query_vector is required (the broker does not embed queries; supply a vector)",
));
}
let _admit = native_admit_on(
svc.channels.as_ref(),
&svc.metrics,
"embedding",
OperationChannel::Vector,
&tenant_id,
None,
)
.await?;
let runtime = svc.require_runtime()?;
let catalog = svc.require_catalog()?;
let mut context = native_service_context(&metadata, &tenant_id, "");
context.scopes.push("udb:vector:read".to_string());
let top_k = resolve_top_k(req.top_k);
let score_floor = {
let env_floor = retrieve_score_threshold();
let requested = req.score_threshold as f32;
if requested > 0.0 {
requested.max(env_floor)
} else {
env_floor
}
};
let stored = runtime
.native_entity_read_for_service(
"embedding",
&context,
source_read_by_name(&tenant_id, &source_name),
)
.await?
.first()
.map(stored_source_from_json)
.filter(|source| source.status == STATUS_ACTIVE)
.ok_or_else(|| embedding_source_not_found_status("retrieve"))?;
let merged_filter = merge_retrieve_filter(&tenant_id, &req.filter_json)?;
let tenant_filter = crate::runtime::executor_utils::json_to_struct(&merged_filter);
let state = catalog.active_for(&context.project_id);
let manifest = &state.manifest;
let collection = stored.collection();
let remaining = remaining_before_deadline(deadline, Instant::now())?;
let has_text = !req.query_text.trim().is_empty();
let result: VectorSet = if has_text {
let search = VectorHybridSearchRequest {
context: None,
collection,
vector: req.query_vector.clone(),
text_query: req.query_text.clone(),
filter: tenant_filter,
limit: top_k,
fusion_weights: retrieve_fusion_weights(),
with_payload: true,
};
let fut = runtime.vector_hybrid_search(manifest, search, context.clone());
match remaining {
Some(budget) => tokio::time::timeout(budget, fut).await.map_err(|_| {
crate::runtime::executor_utils::deadline_exceeded_status(
"embedding",
"retrieve_hybrid_search",
crate::runtime::executor_utils::HTTP_RETRYABLE_BACKOFF_MS,
"retrieve exceeded its deadline during hybrid search",
)
})??,
None => fut.await?,
}
} else {
let search = VectorSearchRequest {
context: None,
collection,
vector: req.query_vector.clone(),
filter: tenant_filter,
limit: top_k,
score_threshold: score_floor,
with_payload: true,
};
let fut = runtime.vector_search(manifest, search, context.clone());
match remaining {
Some(budget) => tokio::time::timeout(budget, fut).await.map_err(|_| {
crate::runtime::executor_utils::deadline_exceeded_status(
"embedding",
"retrieve_vector_search",
crate::runtime::executor_utils::HTTP_RETRYABLE_BACKOFF_MS,
"retrieve exceeded its deadline during vector search",
)
})??,
None => fut.await?,
}
};
let hits = result
.points
.into_iter()
.filter(|point| point.score >= score_floor)
.take(top_k as usize)
.map(|point| embedding_pb::RetrieveHit {
payload_json: retrieve_hit_payload_json(point.payload.as_ref()),
id: point.id,
score: f64::from(point.score),
})
.collect::<Vec<_>>();
Ok(Response::new(embedding_pb::RetrieveResponse {
hits,
message: "ok".to_string(),
error: None,
}))
}