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 std::sync::Arc;
use crate::core::db::PostgresDb; use crate::core::db::PostgresDb;
use crate::core::db::schema::t;
pub fn trace_agent_routes() -> Router<crate::api::types::AppState> { pub fn trace_agent_routes() -> Router<crate::api::types::AppState> {
Router::new() Router::new()
@@ -133,7 +134,10 @@ async fn list_traces_sorted(
{"key": "file_uuid", "match": {"value": file_uuid}} {"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 // Aggregate by trace_id
struct TraceAgg { struct TraceAgg {
@@ -169,25 +173,43 @@ async fn list_traces_sorted(
} }
// Filter by min_faces and sort // 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) .filter(|(_, agg)| agg.face_count >= min_faces)
.map(|(tid, agg)| { .map(|(tid, agg)| {
let duration = (agg.end_frame - agg.start_frame) as f64; 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 }; let avg_conf = if agg.face_count > 0 {
(tid, agg.face_count, agg.start_frame, agg.end_frame, duration, avg_conf) 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(); .collect();
match order_clause { match order_clause {
"face_count DESC" => traces_vec.sort_by(|a, b| b.1.cmp(&a.1)), "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)), _ => traces_vec.sort_by(|a, b| a.2.cmp(&b.2)),
} }
// Apply pagination // Apply pagination
let total_traces = traces_vec.len() as i64; let total_traces = traces_vec.len() as i64;
let total_faces: i64 = points.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 let traces: Vec<TraceInfo> = traces_vec
.into_iter() .into_iter()
@@ -297,12 +319,19 @@ async fn list_trace_faces(
{"key": "trace_id", "match": {"value": trace_id}} {"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; let total_detected: i64 = points.len() as i64;
// Apply pagination // 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(); let mut faces: Vec<TraceFaceItem> = Vec::new();
@@ -440,8 +469,8 @@ async fn select_rep_face<F, T>(
where where
F: Fn(anyhow::Error) -> T, F: Fn(anyhow::Error) -> T,
{ {
use crate::core::db::schema;
use crate::core::db::qdrant_db::QdrantDb; use crate::core::db::qdrant_db::QdrantDb;
use crate::core::db::schema;
use serde_json::json; use serde_json::json;
let video_table = schema::table_name("videos"); let video_table = schema::table_name("videos");
@@ -463,7 +492,10 @@ where
{"key": "trace_id", "match": {"value": trace_id}} {"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,); let face_count: (i64,) = (points.len() as i64,);
struct Candidate { struct Candidate {
@@ -477,14 +509,17 @@ where
} }
// Get top faces by quality from Qdrant // Get top faces by quality from Qdrant
let mut candidates: Vec<Candidate> = points.iter() let mut candidates: Vec<Candidate> = points
.iter()
.filter_map(|p| { .filter_map(|p| {
let payload = &p["payload"]; let payload = &p["payload"];
let bbox = &payload["bbox"]; let bbox = &payload["bbox"];
let w = bbox["width"].as_f64()? as i32; let w = bbox["width"].as_f64()? as i32;
let h = bbox["height"].as_f64()? as i32; let h = bbox["height"].as_f64()? as i32;
let conf = payload["confidence"].as_f64()?; 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; let score = (w as f64 * h as f64) * conf;
Some(Candidate { Some(Candidate {
frame: payload["frame"].as_i64().unwrap_or(0), frame: payload["frame"].as_i64().unwrap_or(0),
@@ -497,7 +532,11 @@ where
}) })
}) })
.collect(); .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(); let rows: Vec<_> = candidates.into_iter().take(10).collect();
if rows.is_empty() { if rows.is_empty() {
@@ -785,8 +824,8 @@ async fn get_cooccurrence(
State(state): State<crate::api::types::AppState>, State(state): State<crate::api::types::AppState>,
Path((file_uuid, identity_uuid_a, identity_uuid_b)): Path<(String, String, String)>, Path((file_uuid, identity_uuid_a, identity_uuid_b)): Path<(String, String, String)>,
) -> Result<Json<CoOccurResponse>, (StatusCode, Json<serde_json::Value>)> { ) -> Result<Json<CoOccurResponse>, (StatusCode, Json<serde_json::Value>)> {
use crate::core::db::schema;
use crate::core::db::qdrant_db::QdrantDb; use crate::core::db::qdrant_db::QdrantDb;
use crate::core::db::schema;
use serde_json::json; use serde_json::json;
let id_table = schema::table_name("identities"); let id_table = schema::table_name("identities");
@@ -841,8 +880,12 @@ async fn get_cooccurrence(
{"key": "identity_id", "match": {"value": id_a.0}} {"key": "identity_id", "match": {"value": id_a.0}}
] ]
}); });
let points_a = qdrant.scroll_all_points("_faces", filter_a, 1000).await.unwrap_or_default(); let points_a = qdrant
let frames_a: std::collections::HashSet<i64> = points_a.iter() .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()) .filter_map(|p| p["payload"]["frame"].as_i64())
.collect(); .collect();
@@ -853,8 +896,12 @@ async fn get_cooccurrence(
{"key": "identity_id", "match": {"value": id_b.0}} {"key": "identity_id", "match": {"value": id_b.0}}
] ]
}); });
let points_b = qdrant.scroll_all_points("_faces", filter_b, 1000).await.unwrap_or_default(); let points_b = qdrant
let cooccur: Option<(i64,)> = points_b.iter() .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()) .filter_map(|p| p["payload"]["frame"].as_i64())
.find(|f| frames_a.contains(f)) .find(|f| frames_a.contains(f))
.map(|f| (f,)); .map(|f| (f,));
@@ -881,12 +928,14 @@ async fn get_cooccurrence(
.unwrap_or(25.0); .unwrap_or(25.0);
// Stage 3: Get trace_ids for both at this frame (from Qdrant _faces) // 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)) .find(|p| p["payload"]["frame"].as_i64() == Some(first_frame))
.and_then(|p| p["payload"]["trace_id"].as_i64()) .and_then(|p| p["payload"]["trace_id"].as_i64())
.map(|t| (t as i32,)); .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)) .find(|p| p["payload"]["frame"].as_i64() == Some(first_frame))
.and_then(|p| p["payload"]["trace_id"].as_i64()) .and_then(|p| p["payload"]["trace_id"].as_i64())
.map(|t| (t as i32,)); .map(|t| (t as i32,));
@@ -941,10 +990,12 @@ async fn get_cooccurrence(
}; };
// Total co-occurrence frames (from Qdrant _faces) // 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()) .filter_map(|p| p["payload"]["frame"].as_i64())
.collect(); .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_map(|p| p["payload"]["frame"].as_i64())
.filter(|f| frames_b.contains(f)) .filter(|f| frames_b.contains(f))
.count() as i64; .count() as i64;
@@ -1006,7 +1057,9 @@ async fn rebuild_tkg(
// Always trigger Rule 2 (even with 0 edges) // Always trigger Rule 2 (even with 0 edges)
info!( info!(
"[TKG] Rebuild completed for {}: {} nodes, {} edges", "[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 { match ingest_rule2(db.pool(), &file_uuid, None, None).await {
Ok(count) => info!("[TKG] Rule 2 created {} relationship chunks", count), 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>> { ) -> Json<Vec<crate::core::tkg::models::TkgOperationLog>> {
use crate::core::tkg::TkgService; use crate::core::tkg::TkgService;
let tkg_service = TkgService::new(state.db); 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) Json(operations)
} }
@@ -1122,9 +1178,13 @@ async fn get_stranger_representative_face(
{"key": "stranger_id", "match": {"value": stranger_id}} {"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()) .and_then(|p| p["payload"]["trace_id"].as_i64())
.map(|t| t as i32) .map(|t| t as i32)
.ok_or(( .ok_or((
@@ -1149,9 +1209,13 @@ async fn get_stranger_thumbnail(
{"key": "stranger_id", "match": {"value": stranger_id}} {"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()) .and_then(|p| p["payload"]["trace_id"].as_i64())
.map(|t| t as i32) .map(|t| t as i32)
.ok_or(( .ok_or((
@@ -1164,15 +1228,6 @@ async fn get_stranger_thumbnail(
// ── TKG Node/Edge Query APIs ───────────────────────────────────── // ── 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)] #[derive(Debug, Deserialize)]
struct TkgNodesRequest { struct TkgNodesRequest {
node_type: Option<String>, node_type: Option<String>,
+10 -25
View File
@@ -1,30 +1,15 @@
use crate::core::config::DATABASE_SCHEMA; /// Schema-qualified table name helper
use once_cell::sync::Lazy; /// 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(|| { /// Alias for t() - used by postgres_db.rs
let schema = DATABASE_SCHEMA.as_str(); pub fn table_name(name: &str) -> String {
let schema = std::env::var("DATABASE_SCHEMA").unwrap_or_else(|_| "dev".to_string());
if schema == "public" { if schema == "public" {
String::new() name.to_string()
} else { } else {
format!("{}.", schema) format!("{}.{}", schema, name)
}
});
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");
} }
} }
+1 -8
View File
@@ -6,6 +6,7 @@ use std::path::Path;
use crate::core::db::postgres_db::PostgresDb; use crate::core::db::postgres_db::PostgresDb;
use crate::core::db::redis_client::RedisClient; use crate::core::db::redis_client::RedisClient;
use crate::core::db::schema::t;
use crate::core::progress::{publish_tkg_progress, TkgPhase, TkgProgress, TkgStats}; use crate::core::progress::{publish_tkg_progress, TkgPhase, TkgProgress, TkgStats};
use std::sync::Arc; use std::sync::Arc;
@@ -121,14 +122,6 @@ async fn build_node_id_map(pool: &PgPool, file_uuid: &str) -> HashMap<(String, S
.collect() .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 ─────────────────────────────────────────────────────── // ── Phase 0: Populate trace_id from face.json ───────────────────────────────────────────────────────