use std::sync::Arc;
use jammi_ai::session::InferenceSession;
use jammi_ai::{AuditError, PerQueryAudit, Session};
use jammi_wire::{attach_audit_detail, parse_query_id, record_from_wire};
use tonic::{Code, Request, Response, Status};
use crate::grpc::proto::audit as pb;
use crate::grpc::proto::audit::audit_service_server::AuditService;
use crate::grpc::wire::{scoped, session_tenant_traced};
pub struct AuditServer {
session: Arc<InferenceSession>,
}
impl AuditServer {
pub fn new(session: Arc<InferenceSession>) -> Self {
Self { session }
}
fn local(&self) -> Session {
Session::new(Arc::clone(&self.session))
}
}
#[tonic::async_trait]
impl AuditService for AuditServer {
#[tracing::instrument(skip(self, request), fields(tenant_id = tracing::field::Empty))]
async fn audit_log(
&self,
request: Request<pb::AuditLogRequest>,
) -> Result<Response<()>, Status> {
let tenant = session_tenant_traced(&request);
let req = request.into_inner();
let records = req
.records
.into_iter()
.map(record_from_proto)
.collect::<Result<Vec<_>, Status>>()?;
let session = self.local();
scoped(&self.session, tenant, || session.audit_log(records))
.await
.map_err(map_audit_error)?;
Ok(Response::new(()))
}
#[tracing::instrument(skip(self, request), fields(tenant_id = tracing::field::Empty))]
async fn audit_fetch_by_query_id(
&self,
request: Request<pb::AuditFetchByQueryIdRequest>,
) -> Result<Response<pb::AuditFetchByQueryIdResponse>, Status> {
let tenant = session_tenant_traced(&request);
let req = request.into_inner();
let query_id = parse_query_id(&req.query_id)?;
let session = self.local();
let record = scoped(&self.session, tenant, || {
session.audit_fetch_by_query_id(query_id)
})
.await
.map_err(map_audit_error)?;
Ok(Response::new(pb::AuditFetchByQueryIdResponse {
record: record.map(Into::into),
}))
}
#[tracing::instrument(skip(self, request), fields(tenant_id = tracing::field::Empty))]
async fn audit_fetch_recent(
&self,
request: Request<pb::AuditFetchRecentRequest>,
) -> Result<Response<pb::AuditFetchRecentResponse>, Status> {
let tenant = session_tenant_traced(&request);
let req = request.into_inner();
let session = self.local();
let records = scoped(&self.session, tenant, || {
session.audit_fetch_recent(req.limit as usize)
})
.await
.map_err(map_audit_error)?;
Ok(Response::new(pb::AuditFetchRecentResponse {
records: records.into_iter().map(Into::into).collect(),
}))
}
#[tracing::instrument(skip(self, request), fields(tenant_id = tracing::field::Empty))]
async fn verify(
&self,
request: Request<pb::VerifyRequest>,
) -> Result<Response<pb::VerifyResponse>, Status> {
let tenant = session_tenant_traced(&request);
let req = request.into_inner();
let record = record_from_wire(
req.record
.ok_or_else(|| Status::invalid_argument("verify requires a record"))?,
)
.map_err(map_audit_error)?;
let session = self.local();
let verified = scoped(&self.session, tenant, || session.audit_verify(record))
.await
.map_err(map_audit_error)?;
Ok(Response::new(pb::VerifyResponse { verified }))
}
}
fn record_from_proto(p: pb::PerQueryAudit) -> Result<PerQueryAudit, Status> {
let query_id = parse_query_id(&p.query_id)?;
let query_lineage: serde_json::Value = if p.query_lineage.is_empty() {
serde_json::Value::Object(serde_json::Map::new())
} else {
serde_json::from_str(&p.query_lineage)
.map_err(|e| Status::invalid_argument(format!("query_lineage is not JSON: {e}")))?
};
PerQueryAudit::new(
query_id,
p.model_id,
p.model_version,
query_lineage,
p.top_k_result_ids,
p.retrieval_scores,
)
.map_err(map_audit_error)
}
fn map_audit_error(err: AuditError) -> Status {
let code = match &err {
AuditError::LengthMismatch { .. } | AuditError::LineageTooLarge { .. } => {
Code::InvalidArgument
}
AuditError::NoTenantBinding | AuditError::MasterKey(_) => Code::FailedPrecondition,
AuditError::SignatureMismatch(_) => Code::DataLoss,
AuditError::Serde(_) | AuditError::Storage(_) | AuditError::Broker(_) => Code::Internal,
};
attach_audit_detail(code, err.to_string(), &err)
}