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:
+92
-37
@@ -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
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
Reference in New Issue
Block a user