refactor(tkg): extract shared t() function to schema module

- Create src/core/db/schema.rs with t() and table_name() functions
- Remove duplicate t() from tkg.rs and trace_agent_api.rs
- postgres_db.rs uses schema::table_name() from shared module
This commit is contained in:
Accusys
2026-07-10 00:08:12 +08:00
parent 5a941f857c
commit 13abb15595
3 changed files with 103 additions and 70 deletions
+92 -37
View File
@@ -10,6 +10,7 @@ use serde::{Deserialize, Serialize};
use std::sync::Arc;
use crate::core::db::PostgresDb;
use crate::core::db::schema::t;
pub fn trace_agent_routes() -> Router<crate::api::types::AppState> {
Router::new()
@@ -133,7 +134,10 @@ async fn list_traces_sorted(
{"key": "file_uuid", "match": {"value": file_uuid}}
]
});
let points = qdrant.scroll_all_points("_faces", face_filter, 2000).await.unwrap_or_default();
let points = qdrant
.scroll_all_points("_faces", face_filter, 2000)
.await
.unwrap_or_default();
// Aggregate by trace_id
struct TraceAgg {
@@ -169,25 +173,43 @@ async fn list_traces_sorted(
}
// Filter by min_faces and sort
let mut traces_vec: Vec<(i32, i64, i64, i64, f64, f64)> = trace_data.into_iter()
let mut traces_vec: Vec<(i32, i64, i64, i64, f64, f64)> = trace_data
.into_iter()
.filter(|(_, agg)| agg.face_count >= min_faces)
.map(|(tid, agg)| {
let duration = (agg.end_frame - agg.start_frame) as f64;
let avg_conf = if agg.face_count > 0 { agg.sum_confidence / agg.face_count as f64 } else { 0.0 };
(tid, agg.face_count, agg.start_frame, agg.end_frame, duration, avg_conf)
let avg_conf = if agg.face_count > 0 {
agg.sum_confidence / agg.face_count as f64
} else {
0.0
};
(
tid,
agg.face_count,
agg.start_frame,
agg.end_frame,
duration,
avg_conf,
)
})
.collect();
match order_clause {
"face_count DESC" => traces_vec.sort_by(|a, b| b.1.cmp(&a.1)),
"duration_sec DESC" => traces_vec.sort_by(|a, b| b.4.partial_cmp(&a.4).unwrap_or(std::cmp::Ordering::Equal)),
"duration_sec DESC" => {
traces_vec.sort_by(|a, b| b.4.partial_cmp(&a.4).unwrap_or(std::cmp::Ordering::Equal))
}
_ => traces_vec.sort_by(|a, b| a.2.cmp(&b.2)),
}
// Apply pagination
let total_traces = traces_vec.len() as i64;
let total_faces: i64 = points.len() as i64;
let traces_vec: Vec<_> = traces_vec.into_iter().skip(db_offset as usize).take(effective_limit as usize).collect();
let traces_vec: Vec<_> = traces_vec
.into_iter()
.skip(db_offset as usize)
.take(effective_limit as usize)
.collect();
let traces: Vec<TraceInfo> = traces_vec
.into_iter()
@@ -297,12 +319,19 @@ async fn list_trace_faces(
{"key": "trace_id", "match": {"value": trace_id}}
]
});
let points = qdrant.scroll_all_points("_faces", trace_filter, 1000).await.unwrap_or_default();
let points = qdrant
.scroll_all_points("_faces", trace_filter, 1000)
.await
.unwrap_or_default();
let total_detected: i64 = points.len() as i64;
// Apply pagination
let paged: Vec<_> = points.into_iter().skip(offset as usize).take(limit as usize).collect();
let paged: Vec<_> = points
.into_iter()
.skip(offset as usize)
.take(limit as usize)
.collect();
let mut faces: Vec<TraceFaceItem> = Vec::new();
@@ -440,8 +469,8 @@ async fn select_rep_face<F, T>(
where
F: Fn(anyhow::Error) -> T,
{
use crate::core::db::schema;
use crate::core::db::qdrant_db::QdrantDb;
use crate::core::db::schema;
use serde_json::json;
let video_table = schema::table_name("videos");
@@ -463,7 +492,10 @@ where
{"key": "trace_id", "match": {"value": trace_id}}
]
});
let points = qdrant.scroll_all_points("_faces", trace_filter, 1000).await.unwrap_or_default();
let points = qdrant
.scroll_all_points("_faces", trace_filter, 1000)
.await
.unwrap_or_default();
let face_count: (i64,) = (points.len() as i64,);
struct Candidate {
@@ -477,14 +509,17 @@ where
}
// Get top faces by quality from Qdrant
let mut candidates: Vec<Candidate> = points.iter()
let mut candidates: Vec<Candidate> = points
.iter()
.filter_map(|p| {
let payload = &p["payload"];
let bbox = &payload["bbox"];
let w = bbox["width"].as_f64()? as i32;
let h = bbox["height"].as_f64()? as i32;
let conf = payload["confidence"].as_f64()?;
if conf <= 0.7 { return None; }
if conf <= 0.7 {
return None;
}
let score = (w as f64 * h as f64) * conf;
Some(Candidate {
frame: payload["frame"].as_i64().unwrap_or(0),
@@ -497,7 +532,11 @@ where
})
})
.collect();
candidates.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal));
candidates.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
let rows: Vec<_> = candidates.into_iter().take(10).collect();
if rows.is_empty() {
@@ -785,8 +824,8 @@ async fn get_cooccurrence(
State(state): State<crate::api::types::AppState>,
Path((file_uuid, identity_uuid_a, identity_uuid_b)): Path<(String, String, String)>,
) -> Result<Json<CoOccurResponse>, (StatusCode, Json<serde_json::Value>)> {
use crate::core::db::schema;
use crate::core::db::qdrant_db::QdrantDb;
use crate::core::db::schema;
use serde_json::json;
let id_table = schema::table_name("identities");
@@ -841,8 +880,12 @@ async fn get_cooccurrence(
{"key": "identity_id", "match": {"value": id_a.0}}
]
});
let points_a = qdrant.scroll_all_points("_faces", filter_a, 1000).await.unwrap_or_default();
let frames_a: std::collections::HashSet<i64> = points_a.iter()
let points_a = qdrant
.scroll_all_points("_faces", filter_a, 1000)
.await
.unwrap_or_default();
let frames_a: std::collections::HashSet<i64> = points_a
.iter()
.filter_map(|p| p["payload"]["frame"].as_i64())
.collect();
@@ -853,8 +896,12 @@ async fn get_cooccurrence(
{"key": "identity_id", "match": {"value": id_b.0}}
]
});
let points_b = qdrant.scroll_all_points("_faces", filter_b, 1000).await.unwrap_or_default();
let cooccur: Option<(i64,)> = points_b.iter()
let points_b = qdrant
.scroll_all_points("_faces", filter_b, 1000)
.await
.unwrap_or_default();
let cooccur: Option<(i64,)> = points_b
.iter()
.filter_map(|p| p["payload"]["frame"].as_i64())
.find(|f| frames_a.contains(f))
.map(|f| (f,));
@@ -881,12 +928,14 @@ async fn get_cooccurrence(
.unwrap_or(25.0);
// Stage 3: Get trace_ids for both at this frame (from Qdrant _faces)
let trace_a: Option<(i32,)> = points_a.iter()
let trace_a: Option<(i32,)> = points_a
.iter()
.find(|p| p["payload"]["frame"].as_i64() == Some(first_frame))
.and_then(|p| p["payload"]["trace_id"].as_i64())
.map(|t| (t as i32,));
let trace_b: Option<(i32,)> = points_b.iter()
let trace_b: Option<(i32,)> = points_b
.iter()
.find(|p| p["payload"]["frame"].as_i64() == Some(first_frame))
.and_then(|p| p["payload"]["trace_id"].as_i64())
.map(|t| (t as i32,));
@@ -941,10 +990,12 @@ async fn get_cooccurrence(
};
// Total co-occurrence frames (from Qdrant _faces)
let frames_b: std::collections::HashSet<i64> = points_b.iter()
let frames_b: std::collections::HashSet<i64> = points_b
.iter()
.filter_map(|p| p["payload"]["frame"].as_i64())
.collect();
let total_cooccurrence_frames: i64 = points_a.iter()
let total_cooccurrence_frames: i64 = points_a
.iter()
.filter_map(|p| p["payload"]["frame"].as_i64())
.filter(|f| frames_b.contains(f))
.count() as i64;
@@ -1006,7 +1057,9 @@ async fn rebuild_tkg(
// Always trigger Rule 2 (even with 0 edges)
info!(
"[TKG] Rebuild completed for {}: {} nodes, {} edges",
file_uuid, r.total_nodes(), r.total_edges()
file_uuid,
r.total_nodes(),
r.total_edges()
);
match ingest_rule2(db.pool(), &file_uuid, None, None).await {
Ok(count) => info!("[TKG] Rule 2 created {} relationship chunks", count),
@@ -1035,7 +1088,10 @@ async fn get_tkg_operations(
) -> Json<Vec<crate::core::tkg::models::TkgOperationLog>> {
use crate::core::tkg::TkgService;
let tkg_service = TkgService::new(state.db);
let operations = tkg_service.get_operations(&file_uuid).await.unwrap_or_default();
let operations = tkg_service
.get_operations(&file_uuid)
.await
.unwrap_or_default();
Json(operations)
}
@@ -1122,9 +1178,13 @@ async fn get_stranger_representative_face(
{"key": "stranger_id", "match": {"value": stranger_id}}
]
});
let points = qdrant.scroll_all_points("_faces", filter, 1).await.unwrap_or_default();
let points = qdrant
.scroll_all_points("_faces", filter, 1)
.await
.unwrap_or_default();
let trace_id: i32 = points.first()
let trace_id: i32 = points
.first()
.and_then(|p| p["payload"]["trace_id"].as_i64())
.map(|t| t as i32)
.ok_or((
@@ -1149,9 +1209,13 @@ async fn get_stranger_thumbnail(
{"key": "stranger_id", "match": {"value": stranger_id}}
]
});
let points = qdrant.scroll_all_points("_faces", filter, 1).await.unwrap_or_default();
let points = qdrant
.scroll_all_points("_faces", filter, 1)
.await
.unwrap_or_default();
let trace_id: i32 = points.first()
let trace_id: i32 = points
.first()
.and_then(|p| p["payload"]["trace_id"].as_i64())
.map(|t| t as i32)
.ok_or((
@@ -1164,15 +1228,6 @@ async fn get_stranger_thumbnail(
// ── TKG Node/Edge Query APIs ─────────────────────────────────────
fn t(name: &str) -> String {
let schema = std::env::var("DATABASE_SCHEMA").unwrap_or_else(|_| "dev".to_string());
if schema == "public" {
name.to_string()
} else {
format!("{}.{}", schema, name)
}
}
#[derive(Debug, Deserialize)]
struct TkgNodesRequest {
node_type: Option<String>,
+10 -25
View File
@@ -1,30 +1,15 @@
use crate::core::config::DATABASE_SCHEMA;
use once_cell::sync::Lazy;
/// Schema-qualified table name helper
/// Returns `name` for public schema, `schema.name` for others
pub fn t(name: &str) -> String {
table_name(name)
}
pub static SCHEMA_PREFIX: Lazy<String> = Lazy::new(|| {
let schema = DATABASE_SCHEMA.as_str();
/// Alias for t() - used by postgres_db.rs
pub fn table_name(name: &str) -> String {
let schema = std::env::var("DATABASE_SCHEMA").unwrap_or_else(|_| "dev".to_string());
if schema == "public" {
String::new()
name.to_string()
} else {
format!("{}.", schema)
}
});
pub fn table_name(table: &str) -> String {
let prefix = SCHEMA_PREFIX.as_str();
if prefix.is_empty() {
table.to_string()
} else {
format!("{}{}", prefix, table)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_table_name_public() {
assert_eq!(table_name("videos"), "videos");
format!("{}.{}", schema, name)
}
}
+1 -8
View File
@@ -6,6 +6,7 @@ use std::path::Path;
use crate::core::db::postgres_db::PostgresDb;
use crate::core::db::redis_client::RedisClient;
use crate::core::db::schema::t;
use crate::core::progress::{publish_tkg_progress, TkgPhase, TkgProgress, TkgStats};
use std::sync::Arc;
@@ -121,14 +122,6 @@ async fn build_node_id_map(pool: &PgPool, file_uuid: &str) -> HashMap<(String, S
.collect()
}
fn t(name: &str) -> String {
let schema = std::env::var("DATABASE_SCHEMA").unwrap_or_else(|_| "dev".to_string());
if schema == "public" {
name.to_string()
} else {
format!("{}.{}", schema, name)
}
}
// ── Phase 0: Populate trace_id from face.json ───────────────────────────────────────────────────────