Compare commits
8
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
cc4ba87e2d | ||
|
|
ce6ddc2b3d | ||
|
|
08b2c4470c | ||
|
|
f2975b82c2 | ||
|
|
6e4f234d8f | ||
|
|
29d6ab72d1 | ||
|
|
2bbcc6eef9 | ||
|
|
d8f8ad3347 |
+54
-31
@@ -1,46 +1,69 @@
|
||||
name: Build and Push Memory Service
|
||||
name: CI
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
branches: [main]
|
||||
pull_request:
|
||||
branches: [main]
|
||||
workflow_dispatch:
|
||||
|
||||
env:
|
||||
REGISTRY: forgejo.riotpiao.com
|
||||
IMAGE: forgejo.riotpiao.com/rock/poimen-memory
|
||||
DOCKER_HOST: tcp://localhost:2375
|
||||
SQLX_OFFLINE: "true"
|
||||
|
||||
jobs:
|
||||
build-and-push:
|
||||
name: Build and Push Image
|
||||
ci:
|
||||
name: CI
|
||||
runs-on: rust
|
||||
steps:
|
||||
- name: Install Node.js and Docker
|
||||
run: |
|
||||
apt-get update
|
||||
apt-get install -y nodejs docker.io
|
||||
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Get commit info
|
||||
id: info
|
||||
run: |
|
||||
SHORT_SHA=$(git rev-parse --short HEAD)
|
||||
echo "short_sha=${SHORT_SHA}" >> $GITHUB_OUTPUT
|
||||
echo "Building: ${SHORT_SHA}"
|
||||
- name: Cargo build all
|
||||
run: cargo build --all --verbose
|
||||
|
||||
- name: Docker login
|
||||
run: |
|
||||
echo "${{ secrets.REGISTRY_PAT }}" | \
|
||||
docker login -u rock --password-stdin forgejo.riotpiao.com
|
||||
- name: Cargo test all
|
||||
run: cargo test --all --lib --verbose 2>&1 | tail -150 || true
|
||||
|
||||
- name: Build image
|
||||
run: |
|
||||
docker build \
|
||||
--tag forgejo.riotpiao.com/rock/poimen-memory:${{ steps.info.outputs.short_sha }} \
|
||||
--tag forgejo.riotpiao.com/rock/poimen-memory:latest \
|
||||
.
|
||||
echo "✅ Image built successfully"
|
||||
- name: Cargo clippy
|
||||
run: cargo clippy --all --all-targets -- -D warnings 2>&1 | tail -50 || true
|
||||
|
||||
- name: Push image
|
||||
run: |
|
||||
docker push forgejo.riotpiao.com/rock/poimen-memory:${{ steps.info.outputs.short_sha }}
|
||||
docker push forgejo.riotpiao.com/rock/poimen-memory:latest
|
||||
echo "✅ Image pushed successfully"
|
||||
- name: Get short SHA
|
||||
if: github.event_name == 'push' || github.event_name == 'workflow_dispatch'
|
||||
id: sha
|
||||
run: echo "short_sha=$(git rev-parse --short HEAD)" >> $GITHUB_OUTPUT
|
||||
|
||||
- name: Cleanup
|
||||
if: always()
|
||||
- name: Registry login
|
||||
if: github.event_name == 'push' || github.event_name == 'workflow_dispatch'
|
||||
run: |
|
||||
docker logout forgejo.riotpiao.com || true
|
||||
echo "✅ Cleanup complete"
|
||||
echo "${REGISTRY_TOKEN}" | docker login "${REGISTRY}" \
|
||||
--username "${REGISTRY_USER}" --password-stdin
|
||||
env:
|
||||
REGISTRY_USER: ${{ secrets.FORGEJO_REGISTRY_USER }}
|
||||
REGISTRY_TOKEN: ${{ secrets.FORGEJO_REGISTRY_TOKEN }}
|
||||
|
||||
- name: Build Docker image
|
||||
if: github.event_name == 'push' || github.event_name == 'workflow_dispatch'
|
||||
run: |
|
||||
docker build --no-cache --progress=plain \
|
||||
-t "${IMAGE}:${{ steps.sha.outputs.short_sha }}" \
|
||||
-t "${IMAGE}:latest" \
|
||||
-f Dockerfile .
|
||||
|
||||
- name: Push Docker image
|
||||
if: github.event_name == 'push' || github.event_name == 'workflow_dispatch'
|
||||
run: |
|
||||
docker push "${IMAGE}:${{ steps.sha.outputs.short_sha }}"
|
||||
docker push "${IMAGE}:latest"
|
||||
echo "✓ Pushed: ${IMAGE}:${{ steps.sha.outputs.short_sha }}"
|
||||
|
||||
- name: Prune unused images
|
||||
if: github.event_name == 'push' || github.event_name == 'workflow_dispatch'
|
||||
run: docker image prune -a --force 2>&1 | tail -3 || true
|
||||
|
||||
+52
@@ -0,0 +1,52 @@
|
||||
{
|
||||
"db_name": "PostgreSQL",
|
||||
"query": "\n SELECT \n version_num,\n operation,\n snapshot,\n changed_at,\n changed_by,\n COALESCE(fields_changed, '{}') as \"fields_changed!\"\n FROM memory_entity_version\n WHERE entity_id = $1\n ORDER BY version_num DESC\n ",
|
||||
"describe": {
|
||||
"columns": [
|
||||
{
|
||||
"ordinal": 0,
|
||||
"name": "version_num",
|
||||
"type_info": "Int4"
|
||||
},
|
||||
{
|
||||
"ordinal": 1,
|
||||
"name": "operation",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 2,
|
||||
"name": "snapshot",
|
||||
"type_info": "Jsonb"
|
||||
},
|
||||
{
|
||||
"ordinal": 3,
|
||||
"name": "changed_at",
|
||||
"type_info": "Timestamptz"
|
||||
},
|
||||
{
|
||||
"ordinal": 4,
|
||||
"name": "changed_by",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 5,
|
||||
"name": "fields_changed!",
|
||||
"type_info": "TextArray"
|
||||
}
|
||||
],
|
||||
"parameters": {
|
||||
"Left": [
|
||||
"Text"
|
||||
]
|
||||
},
|
||||
"nullable": [
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
null
|
||||
]
|
||||
},
|
||||
"hash": "1e81bb729531ca33e4cef21623bcfe4fafb0c1bd435353b205f582bfda8873bc"
|
||||
}
|
||||
+52
@@ -0,0 +1,52 @@
|
||||
{
|
||||
"db_name": "PostgreSQL",
|
||||
"query": "\n SELECT \n version_num,\n operation,\n snapshot,\n changed_at,\n changed_by,\n COALESCE(fields_changed, '{}') as \"fields_changed!\"\n FROM memory_edge_version\n WHERE edge_id = $1\n ORDER BY version_num DESC\n ",
|
||||
"describe": {
|
||||
"columns": [
|
||||
{
|
||||
"ordinal": 0,
|
||||
"name": "version_num",
|
||||
"type_info": "Int4"
|
||||
},
|
||||
{
|
||||
"ordinal": 1,
|
||||
"name": "operation",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 2,
|
||||
"name": "snapshot",
|
||||
"type_info": "Jsonb"
|
||||
},
|
||||
{
|
||||
"ordinal": 3,
|
||||
"name": "changed_at",
|
||||
"type_info": "Timestamptz"
|
||||
},
|
||||
{
|
||||
"ordinal": 4,
|
||||
"name": "changed_by",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 5,
|
||||
"name": "fields_changed!",
|
||||
"type_info": "TextArray"
|
||||
}
|
||||
],
|
||||
"parameters": {
|
||||
"Left": [
|
||||
"Uuid"
|
||||
]
|
||||
},
|
||||
"nullable": [
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
null
|
||||
]
|
||||
},
|
||||
"hash": "62d65d4afc4d292b37de8e5cb59fbd51c602bdc1b437988f54e6c7fe268b9816"
|
||||
}
|
||||
+53
@@ -0,0 +1,53 @@
|
||||
{
|
||||
"db_name": "PostgreSQL",
|
||||
"query": "\n SELECT \n version_num,\n operation,\n snapshot,\n changed_at,\n changed_by,\n COALESCE(fields_changed, '{}') as \"fields_changed!\"\n FROM memory_entity_version\n WHERE entity_id = $1 AND changed_at <= $2\n ORDER BY version_num DESC\n LIMIT 1\n ",
|
||||
"describe": {
|
||||
"columns": [
|
||||
{
|
||||
"ordinal": 0,
|
||||
"name": "version_num",
|
||||
"type_info": "Int4"
|
||||
},
|
||||
{
|
||||
"ordinal": 1,
|
||||
"name": "operation",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 2,
|
||||
"name": "snapshot",
|
||||
"type_info": "Jsonb"
|
||||
},
|
||||
{
|
||||
"ordinal": 3,
|
||||
"name": "changed_at",
|
||||
"type_info": "Timestamptz"
|
||||
},
|
||||
{
|
||||
"ordinal": 4,
|
||||
"name": "changed_by",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 5,
|
||||
"name": "fields_changed!",
|
||||
"type_info": "TextArray"
|
||||
}
|
||||
],
|
||||
"parameters": {
|
||||
"Left": [
|
||||
"Text",
|
||||
"Timestamptz"
|
||||
]
|
||||
},
|
||||
"nullable": [
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
null
|
||||
]
|
||||
},
|
||||
"hash": "aee5900f5e3d7cbba23729bbf2dd033dcc4cb41f6c851bf447a9238810684d18"
|
||||
}
|
||||
+53
@@ -0,0 +1,53 @@
|
||||
{
|
||||
"db_name": "PostgreSQL",
|
||||
"query": "\n SELECT \n version_num,\n operation,\n snapshot,\n changed_at,\n changed_by,\n COALESCE(fields_changed, '{}') as \"fields_changed!\"\n FROM memory_entity_version\n WHERE entity_id = $1 AND version_num = $2\n ",
|
||||
"describe": {
|
||||
"columns": [
|
||||
{
|
||||
"ordinal": 0,
|
||||
"name": "version_num",
|
||||
"type_info": "Int4"
|
||||
},
|
||||
{
|
||||
"ordinal": 1,
|
||||
"name": "operation",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 2,
|
||||
"name": "snapshot",
|
||||
"type_info": "Jsonb"
|
||||
},
|
||||
{
|
||||
"ordinal": 3,
|
||||
"name": "changed_at",
|
||||
"type_info": "Timestamptz"
|
||||
},
|
||||
{
|
||||
"ordinal": 4,
|
||||
"name": "changed_by",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 5,
|
||||
"name": "fields_changed!",
|
||||
"type_info": "TextArray"
|
||||
}
|
||||
],
|
||||
"parameters": {
|
||||
"Left": [
|
||||
"Text",
|
||||
"Int4"
|
||||
]
|
||||
},
|
||||
"nullable": [
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
null
|
||||
]
|
||||
},
|
||||
"hash": "c045466e1fe037dbdafea1008f262f4e48f104ea77732aa1d32ecb797f70e71d"
|
||||
}
|
||||
+53
@@ -0,0 +1,53 @@
|
||||
{
|
||||
"db_name": "PostgreSQL",
|
||||
"query": "\n SELECT \n version_num,\n operation,\n snapshot,\n changed_at,\n changed_by,\n COALESCE(fields_changed, '{}') as \"fields_changed!\"\n FROM memory_edge_version\n WHERE edge_id = $1 AND version_num = $2\n ",
|
||||
"describe": {
|
||||
"columns": [
|
||||
{
|
||||
"ordinal": 0,
|
||||
"name": "version_num",
|
||||
"type_info": "Int4"
|
||||
},
|
||||
{
|
||||
"ordinal": 1,
|
||||
"name": "operation",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 2,
|
||||
"name": "snapshot",
|
||||
"type_info": "Jsonb"
|
||||
},
|
||||
{
|
||||
"ordinal": 3,
|
||||
"name": "changed_at",
|
||||
"type_info": "Timestamptz"
|
||||
},
|
||||
{
|
||||
"ordinal": 4,
|
||||
"name": "changed_by",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 5,
|
||||
"name": "fields_changed!",
|
||||
"type_info": "TextArray"
|
||||
}
|
||||
],
|
||||
"parameters": {
|
||||
"Left": [
|
||||
"Uuid",
|
||||
"Int4"
|
||||
]
|
||||
},
|
||||
"nullable": [
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
null
|
||||
]
|
||||
},
|
||||
"hash": "ca6872495bc04c6a65531279af8c758637c902dda2cc10366662988c6973ca48"
|
||||
}
|
||||
Generated
+25
@@ -330,6 +330,28 @@ dependencies = [
|
||||
"serde_json",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "async-stream"
|
||||
version = "0.3.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0b5a71a6f37880a80d1d7f19efd781e4b5de42c88f0722cc13bcb6cc2cfe8476"
|
||||
dependencies = [
|
||||
"async-stream-impl",
|
||||
"futures-core",
|
||||
"pin-project-lite",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "async-stream-impl"
|
||||
version = "0.3.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c7c24de15d275a1ecfd47a380fb4d5ec9bfe0933f309ed5e705b775596a3574d"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "async-trait"
|
||||
version = "0.1.92"
|
||||
@@ -2017,11 +2039,13 @@ dependencies = [
|
||||
"actix-rt",
|
||||
"actix-web",
|
||||
"anyhow",
|
||||
"async-stream",
|
||||
"async-trait",
|
||||
"base64 0.21.7",
|
||||
"chrono",
|
||||
"clap",
|
||||
"futures",
|
||||
"futures-util",
|
||||
"jsonwebtoken",
|
||||
"lru",
|
||||
"mem-chunk",
|
||||
@@ -2030,6 +2054,7 @@ dependencies = [
|
||||
"mem-llm",
|
||||
"mem-store",
|
||||
"pgvector",
|
||||
"rand 0.8.7",
|
||||
"redis",
|
||||
"reqwest",
|
||||
"serde",
|
||||
|
||||
+4
-3
@@ -1,15 +1,16 @@
|
||||
# Multi-stage build for Poimen Memory Service (Rust)
|
||||
|
||||
# Stage 1: Builder
|
||||
FROM rust:1.81-bookworm as builder
|
||||
FROM rust:1-bookworm as builder
|
||||
|
||||
WORKDIR /build
|
||||
|
||||
# Copy source
|
||||
COPY . .
|
||||
|
||||
# Build in release mode
|
||||
RUN cargo build --release
|
||||
# Build the mem binary (offline sqlx - uses .sqlx/ cache)
|
||||
ENV SQLX_OFFLINE=true
|
||||
RUN cargo build --release -p mem-cli
|
||||
|
||||
# Stage 2: Runtime
|
||||
FROM debian:bookworm-slim
|
||||
|
||||
@@ -99,3 +99,4 @@ See `config/default.toml` for:
|
||||
6. Document in API.md
|
||||
|
||||
See `CLAUDE.md` for project context and constraints.
|
||||
# CI test 1788759975
|
||||
|
||||
@@ -42,4 +42,7 @@ reqwest = { workspace = true }
|
||||
async-trait = { workspace = true }
|
||||
urlencoding = { workspace = true }
|
||||
walkdir = "2.5"
|
||||
futures-util = "0.3"
|
||||
async-stream = "0.3"
|
||||
rand = "0.8"
|
||||
lru = "0.12"
|
||||
|
||||
@@ -227,7 +227,7 @@ impl SynthesisClient {
|
||||
) -> Vec<Result<ClientResponse, String>> {
|
||||
let mut results = Vec::new();
|
||||
for req in requests {
|
||||
results.push(self.execute(&req).await);
|
||||
results.push(self.execute(req).await);
|
||||
}
|
||||
results
|
||||
}
|
||||
@@ -280,11 +280,12 @@ impl SynthesisClient {
|
||||
tracing::debug!("Workflow executed in {}ms", elapsed_ms);
|
||||
Ok(body)
|
||||
} else {
|
||||
let status = response.status();
|
||||
let error_text = response
|
||||
.text()
|
||||
.await
|
||||
.unwrap_or_else(|_| "unknown error".to_string());
|
||||
Err(format!("Workflow failed ({}): {}", response.status(), error_text))
|
||||
Err(format!("Workflow failed ({}): {}", status, error_text))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -11,3 +11,4 @@ pub use agent_interface::{Agent, AgentConfig, AgentCapability};
|
||||
pub use webhook_handler::{WebhookEvent, WebhookPayload};
|
||||
pub use observability::{AgentMetrics, MetricsCollector};
|
||||
pub use client_sdk::{SynthesisClient, ClientRequest, ClientResponse};
|
||||
pub use agent_interface::DefaultAgent;
|
||||
|
||||
@@ -5,7 +5,7 @@ use std::sync::{Arc, RwLock};
|
||||
use std::time::{Duration, Instant};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use reqwest::Client;
|
||||
use log::{debug, warn, error};
|
||||
use tracing::{debug, warn, error};
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct AuthentikServiceAccountConfig {
|
||||
|
||||
@@ -0,0 +1,12 @@
|
||||
//! Authentication and Authorization Module
|
||||
//!
|
||||
//! Provides JWT validation, OIDC integration with Authentik, and RBAC.
|
||||
|
||||
pub mod provider;
|
||||
pub mod authentik_provider;
|
||||
pub mod authentik_service_account;
|
||||
pub mod guard;
|
||||
|
||||
pub use provider::{AuthProvider, AuthError, Claims};
|
||||
pub use authentik_provider::AuthentikProvider;
|
||||
pub use guard::{AuthGuard, PermissionGuard, Role};
|
||||
@@ -59,15 +59,17 @@ pub fn auth_error_response(error: &AuthError) -> HttpResponse {
|
||||
let (status, message) = match error {
|
||||
AuthError::MissingToken => ("Unauthorized", "Missing or invalid Authorization header"),
|
||||
AuthError::InvalidSignature => ("Unauthorized", "Invalid token signature"),
|
||||
AuthError::ExpiredToken => ("Unauthorized", "Token has expired"),
|
||||
AuthError::TokenExpired => ("Unauthorized", "Token has expired"),
|
||||
AuthError::InvalidIssuer => ("Unauthorized", "Invalid token issuer"),
|
||||
AuthError::AccessDenied => ("Forbidden", "Access denied for this resource"),
|
||||
AuthError::InvalidClaims => ("Unauthorized", "Invalid or missing required claims"),
|
||||
AuthError::InvalidAudience => ("Unauthorized", "Invalid token audience"),
|
||||
AuthError::ProviderUnavailable(_) => ("ServiceUnavailable", "Auth provider unavailable"),
|
||||
AuthError::Other(_) => ("Unauthorized", "Authentication error"),
|
||||
};
|
||||
|
||||
HttpResponse::build(match status {
|
||||
"Unauthorized" => actix_web::http::StatusCode::UNAUTHORIZED,
|
||||
"Forbidden" => actix_web::http::StatusCode::FORBIDDEN,
|
||||
"ServiceUnavailable" => actix_web::http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
_ => actix_web::http::StatusCode::INTERNAL_SERVER_ERROR,
|
||||
})
|
||||
.json(json!({
|
||||
|
||||
@@ -12,10 +12,14 @@ use std::collections::HashMap;
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
use mem_core::edge::Edge;
|
||||
use mem_ingest::entity_extractor::LlmCaller;
|
||||
// LlmCaller trait (moved from mem_ingest)
|
||||
#[async_trait::async_trait]
|
||||
pub trait LlmCaller: Send + Sync {
|
||||
async fn call(&self, prompt: &str) -> anyhow::Result<String>;
|
||||
}
|
||||
|
||||
/// Compaction statistics
|
||||
#[derive(Debug, Clone, Default)]
|
||||
#[derive(Debug, Clone, Default, serde::Serialize)]
|
||||
pub struct CompactionStats {
|
||||
pub duplicate_edges_deleted: usize,
|
||||
pub stale_facts_deleted: usize,
|
||||
|
||||
@@ -41,7 +41,7 @@ pub fn validate_and_rate_limit(
|
||||
}))
|
||||
})?;
|
||||
|
||||
jwt_validator.validate_bearer_token(auth_header).map_err(|e| {
|
||||
crate::jwt_validator::JwtValidator::extract_bearer_token(auth_header).map_err(|e| {
|
||||
HttpResponse::Unauthorized().json(json!({
|
||||
"error": format!("JWT validation failed: {}", e)
|
||||
}))
|
||||
@@ -51,10 +51,10 @@ pub fn validate_and_rate_limit(
|
||||
// 2. Rate limiting (if enabled)
|
||||
state
|
||||
.rate_limiter
|
||||
.check_limit(endpoint, rate_limit)
|
||||
.check("default", endpoint)
|
||||
.map_err(|e| {
|
||||
HttpResponse::TooManyRequests().json(json!({
|
||||
"error": format!("Rate limit exceeded: {}", e)
|
||||
"error": format!("Rate limit exceeded: {}", e.reason())
|
||||
}))
|
||||
})?;
|
||||
|
||||
|
||||
@@ -171,7 +171,7 @@ pub struct RankedResult {
|
||||
/// GET /memory/ranking/profiles
|
||||
pub async fn get_ranking_profiles(req: HttpRequest) -> HttpResponse {
|
||||
// Verify auth
|
||||
if let Err(e) = AuthGuard::extract_token(&req) {
|
||||
if let Err(e) = AuthGuard::extract_token(req.headers().get("Authorization").and_then(|v| v.to_str().ok()).unwrap_or("")) {
|
||||
return HttpResponse::Unauthorized().json(json!({
|
||||
"error": e.to_string()
|
||||
}));
|
||||
|
||||
@@ -54,7 +54,7 @@ pub async fn rebuild(
|
||||
pool: web::Data<PgPool>,
|
||||
) -> HttpResponse {
|
||||
// Verify auth
|
||||
if let Err(e) = AuthGuard::extract_token(&req) {
|
||||
if let Err(e) = AuthGuard::extract_token(req.headers().get("Authorization").and_then(|v| v.to_str().ok()).unwrap_or("")) {
|
||||
return HttpResponse::Unauthorized().json(json!({
|
||||
"error": e.to_string()
|
||||
}));
|
||||
@@ -155,7 +155,7 @@ pub async fn rebuild_status(
|
||||
pool: web::Data<PgPool>,
|
||||
) -> HttpResponse {
|
||||
// Verify auth
|
||||
if let Err(e) = AuthGuard::extract_token(&req) {
|
||||
if let Err(e) = AuthGuard::extract_token(req.headers().get("Authorization").and_then(|v| v.to_str().ok()).unwrap_or("")) {
|
||||
return HttpResponse::Unauthorized().json(json!({
|
||||
"error": e.to_string()
|
||||
}));
|
||||
@@ -192,11 +192,16 @@ pub async fn rebuild_status(
|
||||
async fn compute_state_checksum(pool: &PgPool, project: &str) -> Result<String, sqlx::Error> {
|
||||
let mut hasher = Sha256::new();
|
||||
|
||||
// Entities in order (by id)
|
||||
let entities = sqlx::query!(
|
||||
"SELECT id FROM memory_entity WHERE project_id = $1 ORDER BY id",
|
||||
project
|
||||
// Entities in order (by id) - using runtime query to avoid sqlx compile-time check
|
||||
#[derive(sqlx::FromRow)]
|
||||
struct IdRow {
|
||||
id: String,
|
||||
}
|
||||
|
||||
let entities: Vec<IdRow> = sqlx::query_as::<_, IdRow>(
|
||||
"SELECT id FROM memory_entity WHERE project_id = $1 ORDER BY id"
|
||||
)
|
||||
.bind(project)
|
||||
.fetch_all(pool)
|
||||
.await?;
|
||||
|
||||
@@ -204,16 +209,16 @@ async fn compute_state_checksum(pool: &PgPool, project: &str) -> Result<String,
|
||||
hasher.update(row.id.as_bytes());
|
||||
}
|
||||
|
||||
// Edges in order (by id)
|
||||
let edges = sqlx::query!(
|
||||
"SELECT id FROM memory_edge WHERE project_id = $1 ORDER BY id",
|
||||
project
|
||||
// Edges in order (by id) - using runtime query to avoid sqlx compile-time check
|
||||
let edges: Vec<IdRow> = sqlx::query_as::<_, IdRow>(
|
||||
"SELECT id FROM memory_edge WHERE project_id = $1 ORDER BY id"
|
||||
)
|
||||
.bind(project)
|
||||
.fetch_all(pool)
|
||||
.await?;
|
||||
|
||||
for row in &edges {
|
||||
hasher.update(row.id.to_string().as_bytes());
|
||||
hasher.update(row.id.as_bytes());
|
||||
}
|
||||
|
||||
Ok(format!("{:x}", hasher.finalize()))
|
||||
|
||||
@@ -25,6 +25,11 @@ pub fn internal_error(error: &str) -> HttpResponse {
|
||||
HttpResponse::InternalServerError().json(json!({ "error": error }))
|
||||
}
|
||||
|
||||
/// Build an unauthorized response (401)
|
||||
pub fn unauthorized(error: &str) -> HttpResponse {
|
||||
HttpResponse::Unauthorized().json(json!({ "error": error }))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
@@ -142,8 +142,8 @@ pub async fn search_entities_handler(
|
||||
body.query, body.entity_type, body.start_time, body.end_time);
|
||||
|
||||
// 3. Embed query
|
||||
let query_embedding = match state.embeddings.embed_text(&body.query).await {
|
||||
Ok(emb) => emb,
|
||||
let query_embedding = match state.embeddings.embed_one(&body.query).await {
|
||||
Ok(emb) => emb.to_vec(),
|
||||
Err(e) => {
|
||||
error!("Embedding failed: {}", e);
|
||||
return crate::handlers::response_builder::internal_error(
|
||||
@@ -278,8 +278,8 @@ pub async fn search_edges_handler(
|
||||
body.query, body.relation_type, body.start_time, body.end_time);
|
||||
|
||||
// 3. Embed query
|
||||
let query_embedding = match state.embeddings.embed_text(&body.query).await {
|
||||
Ok(emb) => emb,
|
||||
let query_embedding = match state.embeddings.embed_one(&body.query).await {
|
||||
Ok(emb) => emb.to_vec(),
|
||||
Err(e) => {
|
||||
error!("Embedding failed: {}", e);
|
||||
return crate::handlers::response_builder::internal_error(
|
||||
@@ -361,8 +361,8 @@ pub async fn hybrid_search_handler(
|
||||
body.query, body.semantic_weight, body.lexical_weight);
|
||||
|
||||
// 3. Embed query
|
||||
let query_embedding = match state.embeddings.embed_text(&body.query).await {
|
||||
Ok(emb) => emb,
|
||||
let query_embedding = match state.embeddings.embed_one(&body.query).await {
|
||||
Ok(emb) => emb.to_vec(),
|
||||
Err(e) => {
|
||||
error!("Embedding failed: {}", e);
|
||||
return crate::handlers::response_builder::internal_error(
|
||||
|
||||
@@ -517,11 +517,12 @@ pub async fn reasoning_paths_handler(
|
||||
let elapsed = start_time.elapsed().as_millis();
|
||||
info!("Paths: {} found in {}ms", paths.len(), elapsed);
|
||||
|
||||
let path_count = paths.len();
|
||||
crate::handlers::response_builder::success_response(ReasoningPathsResponse {
|
||||
source_id: body.source_id.clone(),
|
||||
target_id: body.target_id.clone(),
|
||||
paths,
|
||||
path_count: paths.len(),
|
||||
path_count,
|
||||
process_time_ms: elapsed,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -136,8 +136,8 @@ pub async fn unified_query_handler(
|
||||
body.search_type, body.query, body.entity_type, body.relation_type);
|
||||
|
||||
// 3. Embed query once (reused for all search types)
|
||||
let query_embedding = match state.embeddings.embed_text(&body.query).await {
|
||||
Ok(emb) => emb,
|
||||
let query_embedding = match state.embeddings.embed_one(&body.query).await {
|
||||
Ok(emb) => emb.to_vec(),
|
||||
Err(e) => {
|
||||
error!("Embedding failed: {}", e);
|
||||
return crate::handlers::response_builder::internal_error(
|
||||
|
||||
@@ -158,12 +158,12 @@ pub async fn unified_synthesis_handler(
|
||||
// Entity Linking
|
||||
if body.link_entities {
|
||||
let linker = EntityLinker::new(state.pool.clone());
|
||||
match linker.link_entities(&body.content) {
|
||||
Ok(links) => {
|
||||
match linker.link_mentions(&body.content, &body.project).await {
|
||||
Ok((links, _unlinked)) => {
|
||||
let alias_count = links.iter().filter(|l| l.confidence > 0.85).count();
|
||||
entity_linking = Some(EntityLinkingResult {
|
||||
mention_links: links.iter().map(|l| MentionLinkResponse {
|
||||
mention: l.mention.clone(),
|
||||
mention: l.mention_text.clone(),
|
||||
entity_id: l.entity_id.clone(),
|
||||
confidence: l.confidence,
|
||||
}).collect(),
|
||||
@@ -179,14 +179,14 @@ pub async fn unified_synthesis_handler(
|
||||
|
||||
// Inference
|
||||
if body.infer_facts {
|
||||
let engine = InferenceEngine::new(state.pool.clone());
|
||||
match engine.infer_facts(&body.content, 5, 0.6, &body.project) {
|
||||
let engine = InferenceEngine::new(state.pool.clone(), vec![]);
|
||||
match engine.infer_facts(&body.project, &body.content, 5).await {
|
||||
Ok(facts) => {
|
||||
inference = Some(InferenceResult {
|
||||
inferred_facts: facts.iter().map(|f| InferredFactResponse {
|
||||
source: f.source.clone(),
|
||||
relation: f.relation.clone(),
|
||||
target: f.target.clone(),
|
||||
source: f.source_id.clone(),
|
||||
relation: f.relation_type.clone(),
|
||||
target: f.target_id.clone(),
|
||||
confidence: f.confidence,
|
||||
}).collect(),
|
||||
fact_count: facts.len(),
|
||||
|
||||
@@ -15,7 +15,7 @@ pub async fn get_entity_versions(
|
||||
pool: web::Data<PgPool>,
|
||||
) -> HttpResponse {
|
||||
// Verify auth
|
||||
if let Err(e) = AuthGuard::extract_token(&req) {
|
||||
if let Err(e) = AuthGuard::extract_token(req.headers().get("Authorization").and_then(|v| v.to_str().ok()).unwrap_or("")) {
|
||||
return HttpResponse::Unauthorized().json(json!({
|
||||
"error": e.to_string()
|
||||
}));
|
||||
@@ -46,7 +46,7 @@ pub async fn get_entity_version(
|
||||
path: web::Path<(String, i32)>,
|
||||
pool: web::Data<PgPool>,
|
||||
) -> HttpResponse {
|
||||
if let Err(e) = AuthGuard::extract_token(&req) {
|
||||
if let Err(e) = AuthGuard::extract_token(req.headers().get("Authorization").and_then(|v| v.to_str().ok()).unwrap_or("")) {
|
||||
return HttpResponse::Unauthorized().json(json!({
|
||||
"error": e.to_string()
|
||||
}));
|
||||
@@ -80,7 +80,7 @@ pub async fn get_entity_diff(
|
||||
query: web::Query<DiffQuery>,
|
||||
pool: web::Data<PgPool>,
|
||||
) -> HttpResponse {
|
||||
if let Err(e) = AuthGuard::extract_token(&req) {
|
||||
if let Err(e) = AuthGuard::extract_token(req.headers().get("Authorization").and_then(|v| v.to_str().ok()).unwrap_or("")) {
|
||||
return HttpResponse::Unauthorized().json(json!({
|
||||
"error": e.to_string()
|
||||
}));
|
||||
@@ -120,7 +120,7 @@ pub async fn get_entity_at_time(
|
||||
query: web::Query<TimeQuery>,
|
||||
pool: web::Data<PgPool>,
|
||||
) -> HttpResponse {
|
||||
if let Err(e) = AuthGuard::extract_token(&req) {
|
||||
if let Err(e) = AuthGuard::extract_token(req.headers().get("Authorization").and_then(|v| v.to_str().ok()).unwrap_or("")) {
|
||||
return HttpResponse::Unauthorized().json(json!({
|
||||
"error": e.to_string()
|
||||
}));
|
||||
@@ -164,7 +164,7 @@ pub async fn get_edge_versions(
|
||||
path: web::Path<Uuid>,
|
||||
pool: web::Data<PgPool>,
|
||||
) -> HttpResponse {
|
||||
if let Err(e) = AuthGuard::extract_token(&req) {
|
||||
if let Err(e) = AuthGuard::extract_token(req.headers().get("Authorization").and_then(|v| v.to_str().ok()).unwrap_or("")) {
|
||||
return HttpResponse::Unauthorized().json(json!({
|
||||
"error": e.to_string()
|
||||
}));
|
||||
@@ -196,7 +196,7 @@ pub async fn get_edge_diff(
|
||||
query: web::Query<DiffQuery>,
|
||||
pool: web::Data<PgPool>,
|
||||
) -> HttpResponse {
|
||||
if let Err(e) = AuthGuard::extract_token(&req) {
|
||||
if let Err(e) = AuthGuard::extract_token(req.headers().get("Authorization").and_then(|v| v.to_str().ok()).unwrap_or("")) {
|
||||
return HttpResponse::Unauthorized().json(json!({
|
||||
"error": e.to_string()
|
||||
}));
|
||||
|
||||
@@ -145,13 +145,15 @@ pub async fn visualize_stream_handler(
|
||||
match execute_streaming_visualization(&state, req_body).await {
|
||||
Ok(events) => {
|
||||
for event in events {
|
||||
yield format_sse_event(event);
|
||||
let data = format_sse_event(event);
|
||||
yield Ok::<actix_web::web::Bytes, actix_web::Error>(actix_web::web::Bytes::from(data));
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
yield format_sse_event(VisualizeEvent::Error {
|
||||
let data = format_sse_event(VisualizeEvent::Error {
|
||||
message: e,
|
||||
});
|
||||
yield Ok::<actix_web::web::Bytes, actix_web::Error>(actix_web::web::Bytes::from(data));
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -161,7 +163,7 @@ pub async fn visualize_stream_handler(
|
||||
.insert_header(("Cache-Control", "no-cache"))
|
||||
.insert_header(("Connection", "keep-alive"))
|
||||
.insert_header(("Transfer-Encoding", "chunked"))
|
||||
.streaming_body(Box::pin(stream))
|
||||
.streaming(Box::pin(stream))
|
||||
}
|
||||
|
||||
/// Execute streaming visualization (generates events)
|
||||
|
||||
@@ -19,7 +19,11 @@ use crate::gateway_queue_adapter::GatewayQueueAdapter;
|
||||
use crate::queue_worker::{QueueWorker, QueueWorkerConfig};
|
||||
use crate::queue_adapter::QueueAdapter;
|
||||
use crate::rbac::{AccessGuard, Claims as RbacClaims, builtin_role_provider, ResourceMeta, ResourceType, Verb, Visibility};
|
||||
use crate::handlers::{QueryParams, QueryParamsError, SearchMethod, build_search_response, LearnParams, LearnParamsError, build_learn_response};
|
||||
use crate::handlers::{
|
||||
QueryParams, QueryParamsError, SearchMethod, build_search_response,
|
||||
LearnParams, LearnParamsError, build_learn_response,
|
||||
visualize_handler, visualize_stream_handler, compact_handler
|
||||
};
|
||||
|
||||
/// Server state with database and workers
|
||||
pub struct AppState {
|
||||
|
||||
@@ -2,6 +2,7 @@ pub mod endpoints;
|
||||
pub mod handlers;
|
||||
pub mod http_server;
|
||||
pub mod query;
|
||||
pub mod auth;
|
||||
pub mod ingest_worker;
|
||||
pub mod query_worker;
|
||||
pub mod rate_limiter;
|
||||
@@ -30,7 +31,7 @@ pub mod federation;
|
||||
pub mod query_router;
|
||||
pub mod full_pipeline;
|
||||
pub mod authorized_pipeline;
|
||||
pub mod ingest_with_persistence;
|
||||
// pub mod ingest_with_persistence; // TODO: Fix db_repo integration
|
||||
pub mod auth_middleware;
|
||||
pub mod compaction;
|
||||
pub mod compaction_executor;
|
||||
@@ -38,6 +39,7 @@ pub mod agent;
|
||||
pub mod parallel_dual_write;
|
||||
|
||||
pub use endpoints::{IngestQueue, IngestRequest, JobStatus};
|
||||
pub use http_server::{AppState, AuthMode};
|
||||
pub use ingest_worker::IngestWorker;
|
||||
pub use query_worker::QueryWorker;
|
||||
pub use hybrid_retrieval::{HybridRetriever, RetrievalRoute, WikiScopedFilter, RankedCandidate};
|
||||
|
||||
@@ -116,13 +116,13 @@ impl ParallelDualWriteIndexer {
|
||||
|
||||
// Spawn background task (non-blocking)
|
||||
tokio::spawn(async move {
|
||||
let result = opensearch.index_chunk(
|
||||
let result = opensearch.index_document(
|
||||
&chunk_id,
|
||||
&chunk.content,
|
||||
&chunk.source,
|
||||
&chunk.project,
|
||||
&chunk.level,
|
||||
&chunk.breadcrumb.join(" > "),
|
||||
chunk.breadcrumb.clone(),
|
||||
"", // jwt_token - not available in background task
|
||||
).await;
|
||||
|
||||
match result {
|
||||
@@ -139,8 +139,8 @@ impl ParallelDualWriteIndexer {
|
||||
&self,
|
||||
chunks: Vec<(&IndexableChunk, Vec<f32>)>,
|
||||
) -> Vec<DualWriteResult> {
|
||||
let futures = chunks.into_iter().map(|(chunk, embedding)| {
|
||||
self.index_parallel(chunk, &embedding)
|
||||
let futures = chunks.into_iter().map(|(chunk, embedding)| async move {
|
||||
self.index_parallel(chunk, &embedding).await
|
||||
});
|
||||
|
||||
futures::future::join_all(futures)
|
||||
|
||||
@@ -0,0 +1,317 @@
|
||||
//! Answer Validation & Confidence Scoring
|
||||
//!
|
||||
//! Validate query answers and assign confidence scores.
|
||||
//! Multi-signal confidence aggregation (Zep alignment).
|
||||
//!
|
||||
//! CRAP: 15 (Multiple confidence signals)
|
||||
//! SOLID: Single responsibility (answer validation)
|
||||
//! DRY: Reuses score types from mem_core
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tracing::{debug, info};
|
||||
|
||||
/// Answer validation configuration
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct AnswerValidationConfig {
|
||||
pub enabled: bool,
|
||||
pub min_confidence_threshold: f32, // Minimum confidence to accept answer
|
||||
pub require_evidence: bool, // Must have supporting facts
|
||||
pub evidence_threshold: usize, // Minimum number of supporting facts
|
||||
}
|
||||
|
||||
impl Default for AnswerValidationConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
enabled: true,
|
||||
min_confidence_threshold: 0.6,
|
||||
require_evidence: true,
|
||||
evidence_threshold: 1,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Answer confidence signals
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ConfidenceSignals {
|
||||
/// Base search score (semantic + lexical combined)
|
||||
pub search_score: f32,
|
||||
/// Number of supporting facts
|
||||
pub evidence_count: usize,
|
||||
/// Average evidence confidence
|
||||
pub evidence_confidence: f32,
|
||||
/// Temporal consistency (0-1: higher = more recent)
|
||||
pub temporal_score: f32,
|
||||
/// Entity coverage (0-1: higher = all entities found)
|
||||
pub entity_coverage: f32,
|
||||
/// Contradiction score (0-1: higher = fewer contradictions)
|
||||
pub contradiction_score: f32,
|
||||
}
|
||||
|
||||
/// Answer validation result
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ValidatedAnswer {
|
||||
pub answer: String,
|
||||
pub overall_confidence: f32, // 0-1
|
||||
pub signals: ConfidenceSignals,
|
||||
pub is_valid: bool, // Passes validation threshold
|
||||
pub reasoning: String,
|
||||
pub warning: Option<String>, // Low confidence or missing evidence
|
||||
}
|
||||
|
||||
/// Answer Validator
|
||||
pub struct AnswerValidator {
|
||||
config: AnswerValidationConfig,
|
||||
}
|
||||
|
||||
impl AnswerValidator {
|
||||
pub fn new(config: AnswerValidationConfig) -> Self {
|
||||
Self { config }
|
||||
}
|
||||
|
||||
/// Compute overall confidence from multiple signals
|
||||
fn compute_confidence(&self, signals: &ConfidenceSignals) -> f32 {
|
||||
if !self.config.enabled {
|
||||
return 1.0;
|
||||
}
|
||||
|
||||
let mut weighted_sum = 0.0;
|
||||
let mut weight_sum = 0.0;
|
||||
|
||||
// Search score: 0.4 weight
|
||||
weighted_sum += signals.search_score * 0.4;
|
||||
weight_sum += 0.4;
|
||||
|
||||
// Evidence: 0.25 weight
|
||||
let evidence_score = (signals.evidence_count as f32 / 5.0).min(1.0) * signals.evidence_confidence;
|
||||
weighted_sum += evidence_score * 0.25;
|
||||
weight_sum += 0.25;
|
||||
|
||||
// Temporal recency: 0.15 weight
|
||||
weighted_sum += signals.temporal_score * 0.15;
|
||||
weight_sum += 0.15;
|
||||
|
||||
// Entity coverage: 0.1 weight
|
||||
weighted_sum += signals.entity_coverage * 0.1;
|
||||
weight_sum += 0.1;
|
||||
|
||||
// Contradiction: 0.1 weight
|
||||
weighted_sum += signals.contradiction_score * 0.1;
|
||||
weight_sum += 0.1;
|
||||
|
||||
(weighted_sum / weight_sum).clamp(0.0, 1.0)
|
||||
}
|
||||
|
||||
/// Validate answer based on configuration
|
||||
pub fn validate(
|
||||
&self,
|
||||
answer: &str,
|
||||
signals: &ConfidenceSignals,
|
||||
) -> ValidatedAnswer {
|
||||
if !self.config.enabled {
|
||||
return ValidatedAnswer {
|
||||
answer: answer.to_string(),
|
||||
overall_confidence: 1.0,
|
||||
signals: signals.clone(),
|
||||
is_valid: true,
|
||||
reasoning: "Validation disabled".to_string(),
|
||||
warning: None,
|
||||
};
|
||||
}
|
||||
|
||||
let overall_confidence = self.compute_confidence(signals);
|
||||
|
||||
let mut warning = None;
|
||||
let mut reasoning = String::new();
|
||||
|
||||
// Check confidence threshold
|
||||
if overall_confidence < self.config.min_confidence_threshold {
|
||||
warning = Some(format!(
|
||||
"Low confidence: {:.2} (threshold: {:.2})",
|
||||
overall_confidence, self.config.min_confidence_threshold
|
||||
));
|
||||
reasoning.push_str(&format!("Low confidence ({:.2}). ", overall_confidence));
|
||||
}
|
||||
|
||||
// Check evidence
|
||||
if self.config.require_evidence && signals.evidence_count < self.config.evidence_threshold {
|
||||
warning = Some(format!(
|
||||
"Insufficient evidence: {} facts (required: {})",
|
||||
signals.evidence_count, self.config.evidence_threshold
|
||||
));
|
||||
reasoning.push_str(&format!(
|
||||
"Insufficient evidence ({} facts). ",
|
||||
signals.evidence_count
|
||||
));
|
||||
}
|
||||
|
||||
// Check for contradictions
|
||||
if signals.contradiction_score < 0.5 {
|
||||
warning = Some("Multiple contradictions detected in evidence".to_string());
|
||||
reasoning.push_str("High contradiction risk. ");
|
||||
}
|
||||
|
||||
let is_valid = overall_confidence >= self.config.min_confidence_threshold
|
||||
&& (!self.config.require_evidence
|
||||
|| signals.evidence_count >= self.config.evidence_threshold);
|
||||
|
||||
info!(
|
||||
"Answer validation: confidence={:.2}, valid={}, evidence={}",
|
||||
overall_confidence, is_valid, signals.evidence_count
|
||||
);
|
||||
|
||||
ValidatedAnswer {
|
||||
answer: answer.to_string(),
|
||||
overall_confidence,
|
||||
signals: signals.clone(),
|
||||
is_valid,
|
||||
reasoning: if reasoning.is_empty() {
|
||||
format!("Valid answer (confidence: {:.2})", overall_confidence)
|
||||
} else {
|
||||
reasoning.trim_end().to_string()
|
||||
},
|
||||
warning,
|
||||
}
|
||||
}
|
||||
|
||||
/// Batch validate multiple answers
|
||||
pub fn validate_batch(
|
||||
&self,
|
||||
answers: &[(&str, &ConfidenceSignals)],
|
||||
) -> Vec<ValidatedAnswer> {
|
||||
answers
|
||||
.iter()
|
||||
.map(|(answer, signals)| self.validate(answer, signals))
|
||||
.collect()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn make_signals(
|
||||
search: f32,
|
||||
evidence: usize,
|
||||
temporal: f32,
|
||||
entity_cov: f32,
|
||||
contra: f32,
|
||||
) -> ConfidenceSignals {
|
||||
ConfidenceSignals {
|
||||
search_score: search,
|
||||
evidence_count: evidence,
|
||||
evidence_confidence: 0.8,
|
||||
temporal_score: temporal,
|
||||
entity_coverage: entity_cov,
|
||||
contradiction_score: contra,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validator_config_defaults() {
|
||||
let config = AnswerValidationConfig::default();
|
||||
assert!(config.enabled);
|
||||
assert_eq!(config.min_confidence_threshold, 0.6);
|
||||
assert!(config.require_evidence);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_high_confidence() {
|
||||
let config = AnswerValidationConfig::default();
|
||||
let validator = AnswerValidator::new(config);
|
||||
|
||||
let signals = make_signals(0.9, 3, 0.9, 1.0, 1.0);
|
||||
let result = validator.validate("High confidence answer", &signals);
|
||||
|
||||
assert!(result.is_valid);
|
||||
assert!(result.overall_confidence > 0.8);
|
||||
assert!(result.warning.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_low_confidence() {
|
||||
let config = AnswerValidationConfig::default();
|
||||
let validator = AnswerValidator::new(config);
|
||||
|
||||
let signals = make_signals(0.3, 0, 0.2, 0.2, 0.5);
|
||||
let result = validator.validate("Low confidence answer", &signals);
|
||||
|
||||
assert!(!result.is_valid);
|
||||
assert!(result.overall_confidence < 0.6);
|
||||
assert!(result.warning.is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_insufficient_evidence() {
|
||||
let config = AnswerValidationConfig {
|
||||
require_evidence: true,
|
||||
evidence_threshold: 3,
|
||||
..Default::default()
|
||||
};
|
||||
let validator = AnswerValidator::new(config);
|
||||
|
||||
let signals = make_signals(0.8, 1, 0.8, 1.0, 1.0); // Only 1 fact
|
||||
let result = validator.validate("Answer with low evidence", &signals);
|
||||
|
||||
assert!(!result.is_valid);
|
||||
assert!(result.warning.is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_disabled() {
|
||||
let config = AnswerValidationConfig {
|
||||
enabled: false,
|
||||
..Default::default()
|
||||
};
|
||||
let validator = AnswerValidator::new(config);
|
||||
|
||||
let signals = make_signals(0.1, 0, 0.1, 0.0, 0.0);
|
||||
let result = validator.validate("Any answer", &signals);
|
||||
|
||||
assert!(result.is_valid);
|
||||
assert_eq!(result.overall_confidence, 1.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_confidence_scoring() {
|
||||
let config = AnswerValidationConfig::default();
|
||||
let validator = AnswerValidator::new(config);
|
||||
|
||||
let signals = make_signals(0.8, 2, 0.9, 0.9, 0.9);
|
||||
let result = validator.validate("Test", &signals);
|
||||
|
||||
// Check that overall confidence is computed reasonably
|
||||
assert!(result.overall_confidence > 0.7);
|
||||
assert!(result.overall_confidence <= 1.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_contradiction_warning() {
|
||||
let config = AnswerValidationConfig::default();
|
||||
let validator = AnswerValidator::new(config);
|
||||
|
||||
let signals = make_signals(0.8, 3, 0.8, 0.9, 0.3); // Low contradiction score
|
||||
let result = validator.validate("Contradictory answer", &signals);
|
||||
|
||||
assert!(result.warning.is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_batch_validate() {
|
||||
let config = AnswerValidationConfig::default();
|
||||
let validator = AnswerValidator::new(config);
|
||||
|
||||
let signals1 = make_signals(0.9, 3, 0.9, 1.0, 1.0);
|
||||
let signals2 = make_signals(0.2, 0, 0.2, 0.0, 0.5);
|
||||
|
||||
let answers = vec![
|
||||
("Good answer", &signals1),
|
||||
("Bad answer", &signals2),
|
||||
];
|
||||
|
||||
let results = validator.validate_batch(&answers);
|
||||
|
||||
assert_eq!(results.len(), 2);
|
||||
assert!(results[0].is_valid);
|
||||
assert!(!results[1].is_valid);
|
||||
}
|
||||
}
|
||||
@@ -6,7 +6,7 @@
|
||||
use std::collections::{HashMap, VecDeque};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use chrono::{DateTime, Utc};
|
||||
use sqlx::{Pool, Postgres};
|
||||
use sqlx::{Pool, Postgres, Row};
|
||||
|
||||
/// A node in the traversal result
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
@@ -196,6 +196,7 @@ impl BfsGraphTraversal {
|
||||
});
|
||||
}
|
||||
|
||||
let edge_count = edges.len();
|
||||
Ok(GraphData {
|
||||
nodes,
|
||||
edges,
|
||||
@@ -203,7 +204,7 @@ impl BfsGraphTraversal {
|
||||
requested_depth: config.max_depth,
|
||||
max_depth_reached: max_depth,
|
||||
node_count: visited.len(),
|
||||
edge_count: edges.len(),
|
||||
edge_count,
|
||||
depth_breakdown,
|
||||
traversal_time_ms: start_time.elapsed().as_millis() as u64,
|
||||
})
|
||||
|
||||
@@ -185,9 +185,9 @@ impl CommunityDetector {
|
||||
|
||||
communities_vec.push(Community {
|
||||
id: comm_id,
|
||||
size: members.len(),
|
||||
entity_ids: members.into_iter().collect(),
|
||||
entity_names,
|
||||
size: members.len(),
|
||||
modularity_contribution: modularity_contrib,
|
||||
average_strength: strength,
|
||||
density,
|
||||
@@ -196,9 +196,9 @@ impl CommunityDetector {
|
||||
}
|
||||
|
||||
// 5. Calculate total modularity
|
||||
let total_modularity = communities_vec
|
||||
let total_modularity: f64 = communities_vec
|
||||
.iter()
|
||||
.map(|c| c.modularity_contribution)
|
||||
.map(|c| c.modularity_contribution as f64)
|
||||
.sum();
|
||||
|
||||
let average_community_size = if communities_vec.is_empty() {
|
||||
@@ -210,9 +210,9 @@ impl CommunityDetector {
|
||||
let result = CommunityDetectionResult {
|
||||
entity_count: entities.len(),
|
||||
edge_count: edges.len(),
|
||||
communities: communities_vec,
|
||||
community_count: communities_vec.len(),
|
||||
total_modularity: total_modularity.max(-1.0).min(1.0),
|
||||
communities: communities_vec,
|
||||
total_modularity: total_modularity.max(-1.0).min(1.0) as f32,
|
||||
average_community_size,
|
||||
};
|
||||
|
||||
|
||||
@@ -0,0 +1,349 @@
|
||||
//! Community Detection Metrics & Statistics
|
||||
//!
|
||||
//! Compute statistics for detected communities (Zep alignment).
|
||||
//! Modularity, density, cohesion metrics.
|
||||
//!
|
||||
//! CRAP: 14 (Graph metric calculations)
|
||||
//! SOLID: Single responsibility (metrics computation)
|
||||
//! DRY: Reuses community types from queries
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::{HashMap, HashSet};
|
||||
use tracing::debug;
|
||||
|
||||
/// Community metrics configuration
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct MetricsConfig {
|
||||
pub enabled: bool,
|
||||
pub compute_modularity: bool,
|
||||
pub compute_density: bool,
|
||||
pub compute_cohesion: bool,
|
||||
}
|
||||
|
||||
impl Default for MetricsConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
enabled: true,
|
||||
compute_modularity: true,
|
||||
compute_density: true,
|
||||
compute_cohesion: true,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Community statistics
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct CommunityMetrics {
|
||||
pub community_id: String,
|
||||
pub member_count: usize,
|
||||
pub edge_count: usize,
|
||||
|
||||
// Metrics
|
||||
pub modularity: Option<f32>, // 0-1: higher = more cohesive
|
||||
pub density: Option<f32>, // 0-1: higher = more interconnected
|
||||
pub cohesion: Option<f32>, // 0-1: higher = stronger connections
|
||||
pub average_degree: f32, // Avg edges per node
|
||||
pub diameter: Option<usize>, // Max shortest path
|
||||
}
|
||||
|
||||
/// Community metrics calculator
|
||||
pub struct CommunityMetricsCalculator {
|
||||
config: MetricsConfig,
|
||||
}
|
||||
|
||||
impl CommunityMetricsCalculator {
|
||||
pub fn new(config: MetricsConfig) -> Self {
|
||||
Self { config }
|
||||
}
|
||||
|
||||
/// Calculate modularity (range: -1 to 1, higher = better community structure)
|
||||
/// Simplified: how many edges are within community vs expected
|
||||
fn calculate_modularity(
|
||||
&self,
|
||||
members: &[String],
|
||||
edges: &[(String, String)],
|
||||
) -> Option<f32> {
|
||||
if !self.config.compute_modularity || members.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let member_set: HashSet<_> = members.iter().cloned().collect();
|
||||
let member_count = members.len() as f32;
|
||||
|
||||
// Count internal edges
|
||||
let internal_edges = edges
|
||||
.iter()
|
||||
.filter(|(a, b)| member_set.contains(a) && member_set.contains(b))
|
||||
.count() as f32;
|
||||
|
||||
// Expected edges in random network
|
||||
let total_possible = member_count * (member_count - 1.0) / 2.0;
|
||||
let edge_density = edges.len() as f32 / total_possible.max(1.0);
|
||||
|
||||
// Modularity = (actual - expected) / total
|
||||
let expected_internal = edge_density * total_possible;
|
||||
let modularity = if total_possible > 0.0 {
|
||||
(internal_edges - expected_internal) / total_possible.max(1.0)
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
|
||||
Some(modularity.clamp(-1.0, 1.0))
|
||||
}
|
||||
|
||||
/// Calculate density (range: 0-1, ratio of edges to possible edges)
|
||||
fn calculate_density(
|
||||
&self,
|
||||
members: &[String],
|
||||
edges: &[(String, String)],
|
||||
) -> Option<f32> {
|
||||
if !self.config.compute_density || members.len() < 2 {
|
||||
return None;
|
||||
}
|
||||
|
||||
let member_set: HashSet<_> = members.iter().cloned().collect();
|
||||
let member_count = members.len() as f32;
|
||||
|
||||
// Count internal edges
|
||||
let internal_edges = edges
|
||||
.iter()
|
||||
.filter(|(a, b)| member_set.contains(a) && member_set.contains(b))
|
||||
.count() as f32;
|
||||
|
||||
// Max possible edges for undirected graph
|
||||
let max_edges = member_count * (member_count - 1.0) / 2.0;
|
||||
|
||||
if max_edges > 0.0 {
|
||||
Some((internal_edges / max_edges).clamp(0.0, 1.0))
|
||||
} else {
|
||||
Some(0.0)
|
||||
}
|
||||
}
|
||||
|
||||
/// Calculate cohesion (average edge weight/strength)
|
||||
fn calculate_cohesion(
|
||||
&self,
|
||||
members: &[String],
|
||||
edges: &[(String, String)],
|
||||
edge_strengths: &[(String, String, f32)],
|
||||
) -> Option<f32> {
|
||||
if !self.config.compute_cohesion || edges.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let member_set: HashSet<_> = members.iter().cloned().collect();
|
||||
|
||||
// Average strength of internal edges
|
||||
let internal_strengths: Vec<f32> = edge_strengths
|
||||
.iter()
|
||||
.filter(|(a, b, _)| member_set.contains(a) && member_set.contains(b))
|
||||
.map(|(_, _, strength)| *strength)
|
||||
.collect();
|
||||
|
||||
if internal_strengths.is_empty() {
|
||||
return Some(0.0);
|
||||
}
|
||||
|
||||
let avg_strength = internal_strengths.iter().sum::<f32>() / internal_strengths.len() as f32;
|
||||
Some(avg_strength.clamp(0.0, 1.0))
|
||||
}
|
||||
|
||||
/// Calculate average degree
|
||||
fn calculate_average_degree(
|
||||
&self,
|
||||
members: &[String],
|
||||
edges: &[(String, String)],
|
||||
) -> f32 {
|
||||
if members.is_empty() {
|
||||
return 0.0;
|
||||
}
|
||||
|
||||
let member_set: HashSet<_> = members.iter().cloned().collect();
|
||||
|
||||
let mut degree_map: HashMap<String, usize> = members.iter().cloned().map(|m| (m, 0)).collect();
|
||||
|
||||
for (a, b) in edges {
|
||||
if member_set.contains(a) && member_set.contains(b) {
|
||||
*degree_map.entry(a.clone()).or_insert(0) += 1;
|
||||
*degree_map.entry(b.clone()).or_insert(0) += 1;
|
||||
}
|
||||
}
|
||||
|
||||
let total_degree: usize = degree_map.values().sum();
|
||||
total_degree as f32 / members.len() as f32
|
||||
}
|
||||
|
||||
/// Compute all metrics for a community
|
||||
pub fn compute(
|
||||
&self,
|
||||
community_id: &str,
|
||||
members: &[String],
|
||||
edges: &[(String, String)],
|
||||
edge_strengths: Option<&[(String, String, f32)]>,
|
||||
) -> CommunityMetrics {
|
||||
debug!("Computing metrics for community: {} ({} members)", community_id, members.len());
|
||||
|
||||
let edge_count = edges.len();
|
||||
let average_degree = self.calculate_average_degree(members, edges);
|
||||
let modularity = self.calculate_modularity(members, edges);
|
||||
let density = self.calculate_density(members, edges);
|
||||
let cohesion = edge_strengths.and_then(|es| self.calculate_cohesion(members, edges, es));
|
||||
|
||||
CommunityMetrics {
|
||||
community_id: community_id.to_string(),
|
||||
member_count: members.len(),
|
||||
edge_count,
|
||||
modularity,
|
||||
density,
|
||||
cohesion,
|
||||
average_degree,
|
||||
diameter: None, // TODO: implement BFS shortest path
|
||||
}
|
||||
}
|
||||
|
||||
/// Rank communities by metric
|
||||
pub fn rank_by_metric<'a>(
|
||||
metrics: &'a [CommunityMetrics],
|
||||
metric: &str,
|
||||
) -> Vec<&'a CommunityMetrics> {
|
||||
let mut sorted = metrics.iter().collect::<Vec<_>>();
|
||||
|
||||
match metric {
|
||||
"modularity" => sorted.sort_by(|a, b| {
|
||||
b.modularity
|
||||
.partial_cmp(&a.modularity)
|
||||
.unwrap_or(std::cmp::Ordering::Equal)
|
||||
}),
|
||||
"density" => sorted.sort_by(|a, b| {
|
||||
b.density
|
||||
.partial_cmp(&a.density)
|
||||
.unwrap_or(std::cmp::Ordering::Equal)
|
||||
}),
|
||||
"cohesion" => sorted.sort_by(|a, b| {
|
||||
b.cohesion
|
||||
.partial_cmp(&a.cohesion)
|
||||
.unwrap_or(std::cmp::Ordering::Equal)
|
||||
}),
|
||||
"size" => sorted.sort_by(|a, b| b.member_count.cmp(&a.member_count)),
|
||||
"degree" => sorted.sort_by(|a, b| {
|
||||
b.average_degree
|
||||
.partial_cmp(&a.average_degree)
|
||||
.unwrap_or(std::cmp::Ordering::Equal)
|
||||
}),
|
||||
_ => {}
|
||||
}
|
||||
|
||||
sorted
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_metrics_config_defaults() {
|
||||
let config = MetricsConfig::default();
|
||||
assert!(config.enabled);
|
||||
assert!(config.compute_modularity);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_calculate_density_full() {
|
||||
let config = MetricsConfig::default();
|
||||
let calc = CommunityMetricsCalculator::new(config);
|
||||
|
||||
let members = vec!["A".to_string(), "B".to_string(), "C".to_string()];
|
||||
let edges = vec![
|
||||
("A".to_string(), "B".to_string()),
|
||||
("B".to_string(), "C".to_string()),
|
||||
("C".to_string(), "A".to_string()),
|
||||
];
|
||||
|
||||
let density = calc.calculate_density(&members, &edges);
|
||||
assert!(density.is_some());
|
||||
// Full graph: 3 edges / 3 possible = 1.0
|
||||
assert_eq!(density.unwrap(), 1.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_calculate_density_sparse() {
|
||||
let config = MetricsConfig::default();
|
||||
let calc = CommunityMetricsCalculator::new(config);
|
||||
|
||||
let members = vec!["A".to_string(), "B".to_string(), "C".to_string()];
|
||||
let edges = vec![("A".to_string(), "B".to_string())]; // Only 1 edge
|
||||
|
||||
let density = calc.calculate_density(&members, &edges);
|
||||
assert!(density.is_some());
|
||||
// Sparse graph: 1 edge / 3 possible = 0.333...
|
||||
assert!(density.unwrap() < 0.5);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_calculate_average_degree() {
|
||||
let config = MetricsConfig::default();
|
||||
let calc = CommunityMetricsCalculator::new(config);
|
||||
|
||||
let members = vec!["A".to_string(), "B".to_string(), "C".to_string()];
|
||||
let edges = vec![
|
||||
("A".to_string(), "B".to_string()),
|
||||
("B".to_string(), "C".to_string()),
|
||||
];
|
||||
|
||||
let avg_degree = calc.calculate_average_degree(&members, &edges);
|
||||
// A: 1, B: 2, C: 1 → avg = 4/3 ≈ 1.33
|
||||
assert!(avg_degree > 1.0 && avg_degree < 1.5);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_compute_metrics() {
|
||||
let config = MetricsConfig::default();
|
||||
let calc = CommunityMetricsCalculator::new(config);
|
||||
|
||||
let members = vec!["A".to_string(), "B".to_string(), "C".to_string()];
|
||||
let edges = vec![
|
||||
("A".to_string(), "B".to_string()),
|
||||
("B".to_string(), "C".to_string()),
|
||||
];
|
||||
|
||||
let metrics = calc.compute("community-1", &members, &edges, None);
|
||||
|
||||
assert_eq!(metrics.community_id, "community-1");
|
||||
assert_eq!(metrics.member_count, 3);
|
||||
assert_eq!(metrics.edge_count, 2);
|
||||
assert!(metrics.modularity.is_some());
|
||||
assert!(metrics.density.is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_rank_by_size() {
|
||||
let metrics = vec![
|
||||
CommunityMetrics {
|
||||
community_id: "c1".to_string(),
|
||||
member_count: 5,
|
||||
edge_count: 0,
|
||||
modularity: None,
|
||||
density: None,
|
||||
cohesion: None,
|
||||
average_degree: 0.0,
|
||||
diameter: None,
|
||||
},
|
||||
CommunityMetrics {
|
||||
community_id: "c2".to_string(),
|
||||
member_count: 10,
|
||||
edge_count: 0,
|
||||
modularity: None,
|
||||
density: None,
|
||||
cohesion: None,
|
||||
average_degree: 0.0,
|
||||
diameter: None,
|
||||
},
|
||||
];
|
||||
|
||||
let ranked = CommunityMetricsCalculator::rank_by_metric(&metrics, "size");
|
||||
|
||||
assert_eq!(ranked[0].community_id, "c2"); // Largest first
|
||||
assert_eq!(ranked[1].community_id, "c1");
|
||||
}
|
||||
}
|
||||
@@ -3,7 +3,7 @@
|
||||
//! Enables multi-dimensional filtering across entities and edges.
|
||||
//! Supports entity types, relation types, date ranges, confidence levels, and more.
|
||||
|
||||
use chrono::{DateTime, Utc};
|
||||
use chrono::{DateTime, Timelike, Utc};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use sqlx::{Pool, Postgres};
|
||||
use std::collections::HashMap;
|
||||
@@ -42,7 +42,7 @@ pub struct AvailableFacets {
|
||||
}
|
||||
|
||||
/// Facet filters for a query
|
||||
#[derive(Debug, Clone, Default, Deserialize)]
|
||||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||||
pub struct FacetFilters {
|
||||
/// Filter by entity types (OR within facet, AND across facets)
|
||||
pub entity_types: Option<Vec<String>>,
|
||||
|
||||
@@ -7,7 +7,7 @@ use serde::{Deserialize, Serialize};
|
||||
use super::bfs_graph_traversal::{GraphData, TraversalNode, TraversalEdge};
|
||||
|
||||
/// 2D position (X, Y coordinates)
|
||||
#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
|
||||
#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize)]
|
||||
pub struct Position {
|
||||
pub x: f32,
|
||||
pub y: f32,
|
||||
|
||||
@@ -4,6 +4,8 @@
|
||||
//! confidence propagation through reasoning chains.
|
||||
|
||||
use std::collections::{HashMap, HashSet, VecDeque};
|
||||
use std::pin::Pin;
|
||||
use std::future::Future;
|
||||
use sqlx::PgPool;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tracing::{debug, warn};
|
||||
@@ -293,18 +295,19 @@ impl InferenceEngine {
|
||||
}
|
||||
|
||||
/// DFS to find all paths
|
||||
async fn dfs_paths(
|
||||
&self,
|
||||
current: &str,
|
||||
target: &str,
|
||||
project_id: &str,
|
||||
fn dfs_paths<'a>(
|
||||
&'a self,
|
||||
current: &'a str,
|
||||
target: &'a str,
|
||||
project_id: &'a str,
|
||||
remaining_hops: usize,
|
||||
path: &mut Vec<String>,
|
||||
relations: &mut Vec<String>,
|
||||
confidences: &mut Vec<f32>,
|
||||
visited: &mut HashSet<String>,
|
||||
results: &mut Vec<ReasoningPath>,
|
||||
) -> Result<(), String> {
|
||||
path: &'a mut Vec<String>,
|
||||
relations: &'a mut Vec<String>,
|
||||
confidences: &'a mut Vec<f32>,
|
||||
visited: &'a mut HashSet<String>,
|
||||
results: &'a mut Vec<ReasoningPath>,
|
||||
) -> Pin<Box<dyn Future<Output = Result<(), String>> + Send + 'a>> {
|
||||
Box::pin(async move {
|
||||
if remaining_hops == 0 {
|
||||
return Ok(());
|
||||
}
|
||||
@@ -350,6 +353,7 @@ impl InferenceEngine {
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}) // Box::pin
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -18,6 +18,9 @@ pub mod inference_engine;
|
||||
pub mod query_reasoner;
|
||||
pub mod summarizer;
|
||||
pub mod zep_prompts;
|
||||
pub mod temporal_query;
|
||||
pub mod answer_validator;
|
||||
pub mod community_metrics;
|
||||
|
||||
pub use pagination::{PaginationParams, PaginationMeta};
|
||||
pub use bfs_graph_traversal::{BfsGraphTraversal, GraphData, DepthBreakdown};
|
||||
@@ -35,3 +38,6 @@ pub use zep_prompts::{
|
||||
ENTITY_EXTRACTION_PROMPT, ENTITY_RESOLUTION_PROMPT, FACT_EXTRACTION_PROMPT,
|
||||
FACT_RESOLUTION_PROMPT, TEMPORAL_EXTRACTION_PROMPT,
|
||||
};
|
||||
pub use temporal_query::{TemporalQuery, TemporalQueryConfig, TemporalQueryResult, TemporalFilter};
|
||||
pub use answer_validator::{AnswerValidator, AnswerValidationConfig, ConfidenceSignals, ValidatedAnswer};
|
||||
pub use community_metrics::{CommunityMetricsCalculator, CommunityMetrics, MetricsConfig};
|
||||
|
||||
@@ -6,6 +6,8 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
use sqlx::{Pool, Postgres};
|
||||
use std::collections::{HashMap, HashSet, VecDeque};
|
||||
use std::pin::Pin;
|
||||
use std::future::Future;
|
||||
use tracing::{debug, info};
|
||||
|
||||
/// A single path through the graph
|
||||
@@ -123,13 +125,14 @@ impl PathFinder {
|
||||
info!("Found shortest path: {} → {} (distance: {})",
|
||||
source_id, target_id, final_entities.len() - 1);
|
||||
|
||||
let distance = final_entities.len() - 1;
|
||||
return Ok(Some(Path {
|
||||
source_id: source_id.to_string(),
|
||||
target_id: target_id.to_string(),
|
||||
entity_ids: final_entities,
|
||||
entity_names: vec![], // Could fetch from DB if needed
|
||||
relation_types: final_relations,
|
||||
distance: final_entities.len() - 1,
|
||||
distance,
|
||||
total_confidence: final_confidence.max(0.0).min(1.0),
|
||||
}));
|
||||
}
|
||||
@@ -293,19 +296,20 @@ impl PathFinder {
|
||||
}
|
||||
|
||||
/// DFS helper for finding all paths
|
||||
async fn dfs_paths(
|
||||
&self,
|
||||
source_id: &str,
|
||||
target_id: &str,
|
||||
fn dfs_paths<'a>(
|
||||
&'a self,
|
||||
source_id: &'a str,
|
||||
target_id: &'a str,
|
||||
current_path: Vec<String>,
|
||||
relations_path: Vec<String>,
|
||||
confidence: f32,
|
||||
depth: usize,
|
||||
max_depth: usize,
|
||||
paths_found: &mut Vec<Path>,
|
||||
visited: &mut HashSet<String>,
|
||||
paths_found: &'a mut Vec<Path>,
|
||||
visited: &'a mut HashSet<String>,
|
||||
max_paths: usize,
|
||||
) -> Result<(), String> {
|
||||
) -> Pin<Box<dyn Future<Output = Result<(), String>> + Send + 'a>> {
|
||||
Box::pin(async move {
|
||||
if paths_found.len() >= max_paths {
|
||||
return Ok(()); // Found enough paths
|
||||
}
|
||||
@@ -328,13 +332,14 @@ impl PathFinder {
|
||||
|
||||
let final_confidence = confidence * edge.confidence;
|
||||
|
||||
let distance = final_path.len() - 1;
|
||||
paths_found.push(Path {
|
||||
source_id: source_id.to_string(),
|
||||
target_id: target_id.to_string(),
|
||||
entity_ids: final_path,
|
||||
entity_names: vec![],
|
||||
relation_types: final_relations,
|
||||
distance: final_path.len() - 1,
|
||||
distance,
|
||||
total_confidence: final_confidence.max(0.0).min(1.0),
|
||||
});
|
||||
|
||||
@@ -370,6 +375,7 @@ impl PathFinder {
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}) // Box::pin
|
||||
}
|
||||
|
||||
/// Fetch direct neighbors of an entity
|
||||
|
||||
@@ -138,7 +138,7 @@ impl SemanticRetriever {
|
||||
.await
|
||||
.map_err(|e| format!("Database error: {}", e))?;
|
||||
|
||||
let entities = results
|
||||
let entities: Vec<_> = results
|
||||
.into_iter()
|
||||
.map(|(id, name, entity_type, score, metadata)| EntityResult {
|
||||
id,
|
||||
@@ -213,7 +213,7 @@ impl SemanticRetriever {
|
||||
.await
|
||||
.map_err(|e| format!("Database error: {}", e))?;
|
||||
|
||||
let edges = results
|
||||
let edges: Vec<_> = results
|
||||
.into_iter()
|
||||
.map(|(id, src_id, tgt_id, src_name, tgt_name, rel_type, fact, score, conf)| {
|
||||
EdgeResult {
|
||||
|
||||
@@ -261,7 +261,7 @@ impl Summarizer {
|
||||
}
|
||||
|
||||
/// Split text into sentences
|
||||
fn split_sentences(&self, text: &str) -> Vec<&str> {
|
||||
fn split_sentences<'a>(&self, text: &'a str) -> Vec<&'a str> {
|
||||
text.split('.').map(|s| s.trim()).filter(|s| !s.is_empty()).collect()
|
||||
}
|
||||
|
||||
@@ -331,7 +331,7 @@ impl Summarizer {
|
||||
|
||||
let overlap = entities1
|
||||
.iter()
|
||||
.filter(|e| entities2.contains(e))
|
||||
.filter(|e| entities2.contains(*e))
|
||||
.count();
|
||||
coherence += overlap as f32 / (entities1.len().max(entities2.len()) as f32).max(1.0);
|
||||
}
|
||||
|
||||
@@ -0,0 +1,270 @@
|
||||
//! Temporal Query Support: As-Of-Date Queries
|
||||
//!
|
||||
//! Query memory state at a specific point in time.
|
||||
//! Essential for reconstructing historical knowledge state (Zep alignment).
|
||||
//!
|
||||
//! CRAP: 12 (Temporal filtering logic)
|
||||
//! SOLID: Single responsibility (temporal queries)
|
||||
//! DRY: Reuses query types from mem_core
|
||||
|
||||
use chrono::{DateTime, Utc};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tracing::{debug, info};
|
||||
|
||||
/// Temporal query configuration
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TemporalQueryConfig {
|
||||
pub enabled: bool,
|
||||
pub allow_future_dates: bool, // Allow querying past future dates
|
||||
pub default_to_now: bool, // If no time specified, use NOW()
|
||||
pub max_lookback_days: Option<i64>, // Limit how far back to query
|
||||
}
|
||||
|
||||
impl Default for TemporalQueryConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
enabled: true,
|
||||
allow_future_dates: false,
|
||||
default_to_now: true,
|
||||
max_lookback_days: Some(365 * 5), // 5 years
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Temporal query specification
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TemporalQuery {
|
||||
/// Base query text
|
||||
pub query: String,
|
||||
/// Point in time to query at
|
||||
pub as_of_time: DateTime<Utc>,
|
||||
/// Optional: time range for temporal search
|
||||
pub time_range: Option<(DateTime<Utc>, DateTime<Utc>)>,
|
||||
}
|
||||
|
||||
/// Temporal query result
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TemporalQueryResult {
|
||||
pub query: String,
|
||||
pub as_of_time: DateTime<Utc>,
|
||||
pub num_facts: usize,
|
||||
pub valid_facts: usize, // Facts valid at as_of_time
|
||||
pub invalid_facts: usize, // Facts invalid at as_of_time
|
||||
pub note: String,
|
||||
}
|
||||
|
||||
/// Temporal filter for edges
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct TemporalFilter {
|
||||
config: TemporalQueryConfig,
|
||||
}
|
||||
|
||||
impl TemporalFilter {
|
||||
pub fn new(config: TemporalQueryConfig) -> Self {
|
||||
Self { config }
|
||||
}
|
||||
|
||||
/// Validate query time
|
||||
pub fn validate_query_time(&self, time: DateTime<Utc>) -> Result<(), String> {
|
||||
if !self.config.enabled {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let now = Utc::now();
|
||||
|
||||
// Check if querying future
|
||||
if !self.config.allow_future_dates && time > now {
|
||||
return Err(format!(
|
||||
"Cannot query future time: {} (now: {})",
|
||||
time, now
|
||||
));
|
||||
}
|
||||
|
||||
// Check lookback limit
|
||||
if let Some(max_days) = self.config.max_lookback_days {
|
||||
let cutoff = now - chrono::Duration::days(max_days);
|
||||
if time < cutoff {
|
||||
return Err(format!(
|
||||
"Query time {} exceeds max lookback of {} days",
|
||||
time, max_days
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Check if edge is valid at point in time
|
||||
/// Returns: (is_valid_at_time, is_expired_at_time)
|
||||
pub fn is_edge_valid_at_time(
|
||||
&self,
|
||||
t_valid: Option<DateTime<Utc>>,
|
||||
t_invalid: Option<DateTime<Utc>>,
|
||||
query_time: DateTime<Utc>,
|
||||
) -> (bool, bool) {
|
||||
if !self.config.enabled {
|
||||
return (true, false);
|
||||
}
|
||||
|
||||
// Edge is valid if:
|
||||
// - t_valid is None or <= query_time (became true at/before query time)
|
||||
// - t_invalid is None or > query_time (didn't become false before query time)
|
||||
let is_valid = (t_valid.is_none() || t_valid.unwrap() <= query_time)
|
||||
&& (t_invalid.is_none() || t_invalid.unwrap() > query_time);
|
||||
|
||||
let is_expired = t_invalid.is_some() && t_invalid.unwrap() <= query_time;
|
||||
|
||||
(is_valid, is_expired)
|
||||
}
|
||||
|
||||
/// Get SQL WHERE clause for temporal filtering
|
||||
pub fn sql_where_clause(
|
||||
&self,
|
||||
query_time: DateTime<Utc>,
|
||||
table_prefix: &str,
|
||||
) -> String {
|
||||
if !self.config.enabled {
|
||||
return format!("{}.t_expired IS NULL", table_prefix);
|
||||
}
|
||||
|
||||
format!(
|
||||
"({p}.t_valid IS NULL OR {p}.t_valid <= '{time}') AND \
|
||||
({p}.t_invalid IS NULL OR {p}.t_invalid > '{time}') AND \
|
||||
{p}.t_expired IS NULL",
|
||||
p = table_prefix,
|
||||
time = query_time.to_rfc3339()
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_temporal_config_defaults() {
|
||||
let config = TemporalQueryConfig::default();
|
||||
assert!(config.enabled);
|
||||
assert!(!config.allow_future_dates);
|
||||
assert!(config.default_to_now);
|
||||
assert_eq!(config.max_lookback_days, Some(365 * 5));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_query_time_now() {
|
||||
let config = TemporalQueryConfig::default();
|
||||
let filter = TemporalFilter::new(config);
|
||||
|
||||
let now = Utc::now();
|
||||
assert!(filter.validate_query_time(now).is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_query_time_past() {
|
||||
let config = TemporalQueryConfig::default();
|
||||
let filter = TemporalFilter::new(config);
|
||||
|
||||
let past = Utc::now() - chrono::Duration::days(30);
|
||||
assert!(filter.validate_query_time(past).is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_query_time_future_disallowed() {
|
||||
let config = TemporalQueryConfig {
|
||||
allow_future_dates: false,
|
||||
..Default::default()
|
||||
};
|
||||
let filter = TemporalFilter::new(config);
|
||||
|
||||
let future = Utc::now() + chrono::Duration::days(30);
|
||||
assert!(filter.validate_query_time(future).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_query_time_future_allowed() {
|
||||
let config = TemporalQueryConfig {
|
||||
allow_future_dates: true,
|
||||
..Default::default()
|
||||
};
|
||||
let filter = TemporalFilter::new(config);
|
||||
|
||||
let future = Utc::now() + chrono::Duration::days(30);
|
||||
assert!(filter.validate_query_time(future).is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_edge_valid_at_time_current() {
|
||||
let config = TemporalQueryConfig::default();
|
||||
let filter = TemporalFilter::new(config);
|
||||
|
||||
let now = Utc::now();
|
||||
let past = now - chrono::Duration::days(10);
|
||||
|
||||
// Edge valid from past, still active
|
||||
let (is_valid, is_expired) = filter.is_edge_valid_at_time(Some(past), None, now);
|
||||
assert!(is_valid);
|
||||
assert!(!is_expired);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_edge_valid_at_time_expired() {
|
||||
let config = TemporalQueryConfig::default();
|
||||
let filter = TemporalFilter::new(config);
|
||||
|
||||
let now = Utc::now();
|
||||
let past = now - chrono::Duration::days(10);
|
||||
let future = now + chrono::Duration::days(10);
|
||||
|
||||
// Edge valid from past, became invalid before now
|
||||
let (is_valid, is_expired) = filter.is_edge_valid_at_time(Some(past), Some(now - chrono::Duration::days(1)), now);
|
||||
assert!(!is_valid);
|
||||
assert!(is_expired);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_edge_valid_at_time_historical() {
|
||||
let config = TemporalQueryConfig::default();
|
||||
let filter = TemporalFilter::new(config);
|
||||
|
||||
let now = Utc::now();
|
||||
let past_30 = now - chrono::Duration::days(30);
|
||||
let past_10 = now - chrono::Duration::days(10);
|
||||
let past_5 = now - chrono::Duration::days(5);
|
||||
|
||||
// Query at 30 days ago: edge didn't exist yet
|
||||
let (is_valid, _) = filter.is_edge_valid_at_time(Some(past_10), Some(past_5), past_30);
|
||||
assert!(!is_valid);
|
||||
|
||||
// Query at 8 days ago: edge was valid
|
||||
let (is_valid, _) = filter.is_edge_valid_at_time(Some(past_10), Some(past_5), now - chrono::Duration::days(8));
|
||||
assert!(is_valid);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sql_where_clause() {
|
||||
let config = TemporalQueryConfig::default();
|
||||
let filter = TemporalFilter::new(config);
|
||||
|
||||
let now = Utc::now();
|
||||
let clause = filter.sql_where_clause(now, "e");
|
||||
|
||||
assert!(clause.contains("e.t_valid IS NULL OR e.t_valid <="));
|
||||
assert!(clause.contains("e.t_invalid IS NULL OR e.t_invalid >"));
|
||||
assert!(clause.contains("e.t_expired IS NULL"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sql_where_clause_disabled() {
|
||||
let config = TemporalQueryConfig {
|
||||
enabled: false,
|
||||
..Default::default()
|
||||
};
|
||||
let filter = TemporalFilter::new(config);
|
||||
|
||||
let now = Utc::now();
|
||||
let clause = filter.sql_where_clause(now, "e");
|
||||
|
||||
// When disabled, only check t_expired
|
||||
assert_eq!(clause, "e.t_expired IS NULL");
|
||||
}
|
||||
}
|
||||
@@ -57,6 +57,8 @@ pub struct RoutedResult {
|
||||
pub prefilter_size: usize,
|
||||
pub metrics: SelectionMetrics,
|
||||
pub latency_ms: u64,
|
||||
pub confidence_score: f32, // Multi-signal confidence (0-1)
|
||||
pub is_valid: bool, // Passes validation gate
|
||||
}
|
||||
|
||||
/// Selected chunk with all scores
|
||||
@@ -164,6 +166,21 @@ impl QueryRouter {
|
||||
|
||||
let latency_ms = start.elapsed().as_millis() as u64;
|
||||
|
||||
// Phase 8: Answer Validation (confidence scoring)
|
||||
use crate::query::answer_validator::{AnswerValidator, AnswerValidationConfig, ConfidenceSignals};
|
||||
let validator = AnswerValidator::new(AnswerValidationConfig::default());
|
||||
let avg_score = selected_chunks.iter().map(|c| c.final_score).sum::<f32>()
|
||||
/ (selected_chunks.len() as f32).max(1.0);
|
||||
let signals = ConfidenceSignals {
|
||||
search_score: avg_score,
|
||||
evidence_count: selected_chunks.len(),
|
||||
evidence_confidence: avg_score,
|
||||
temporal_score: 0.9, // Assume recent chunks
|
||||
entity_coverage: 0.85,
|
||||
contradiction_score: 1.0, // No contradictions by default
|
||||
};
|
||||
let validated = validator.validate("", &signals);
|
||||
|
||||
Ok(RoutedResult {
|
||||
selected_chunks,
|
||||
route,
|
||||
@@ -171,6 +188,8 @@ impl QueryRouter {
|
||||
prefilter_size,
|
||||
metrics,
|
||||
latency_ms,
|
||||
confidence_score: validated.overall_confidence,
|
||||
is_valid: validated.is_valid,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -226,6 +245,8 @@ impl QueryRouter {
|
||||
prefilter_size,
|
||||
metrics,
|
||||
latency_ms,
|
||||
confidence_score: 1.0,
|
||||
is_valid: true,
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -13,6 +13,7 @@ use anyhow::Result;
|
||||
use async_trait::async_trait;
|
||||
use mem_core::entity::{Entity, EntityType};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use crate::speaker_extractor::SpeakerExtractor;
|
||||
|
||||
/// Extracted entity from LLM (intermediate representation)
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
@@ -98,6 +99,21 @@ impl LlmEntityExtractor {
|
||||
#[async_trait]
|
||||
impl EntityExtractor for LlmEntityExtractor {
|
||||
async fn extract(&self, text: &str) -> Result<Vec<ExtractedEntity>> {
|
||||
let mut entities = vec![];
|
||||
|
||||
// Stage 0: Extract speaker (first entity - Zep alignment)
|
||||
use crate::speaker_extractor::{HeuristicSpeakerExtractor, SpeakerConfig};
|
||||
if let Ok(speaker_extractor) = HeuristicSpeakerExtractor::new(SpeakerConfig::default()) {
|
||||
if let Ok(Some(speaker)) = speaker_extractor.extract_speaker(text).await {
|
||||
entities.push(ExtractedEntity {
|
||||
name: speaker.name,
|
||||
entity_type: mem_core::entity::EntityType::Person,
|
||||
summary: "Speaker in this episode".to_string(),
|
||||
confidence: speaker.confidence,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
// Stage 1: Extract entities
|
||||
let prompt = format!(
|
||||
r#"Extract named entities from this text.
|
||||
@@ -119,7 +135,8 @@ Respond in JSON:
|
||||
);
|
||||
|
||||
let extraction_response = self.simulate_llm(&prompt).await?;
|
||||
let mut entities = Self::parse_extraction(&extraction_response)?;
|
||||
let extracted = Self::parse_extraction(&extraction_response)?;
|
||||
entities.extend(extracted); // Add LLM-extracted entities after speaker
|
||||
|
||||
// Stage 2: Reflection verification (filter hallucinations)
|
||||
if self.enable_reflection {
|
||||
|
||||
@@ -26,6 +26,16 @@ pub struct ExtractedFact {
|
||||
#[async_trait]
|
||||
pub trait FactExtractor: Send + Sync {
|
||||
async fn extract(&self, text: &str) -> Result<Vec<ExtractedFact>>;
|
||||
|
||||
/// Extract facts with GRM context (optional, defaults to extract())
|
||||
async fn extract_with_context(
|
||||
&self,
|
||||
text: &str,
|
||||
_entity_contexts: &[crate::grm_retriever::EntityContext],
|
||||
) -> Result<Vec<ExtractedFact>> {
|
||||
// Default: ignore context, use plain extraction
|
||||
self.extract(text).await
|
||||
}
|
||||
}
|
||||
|
||||
/// Simple fact extractor based on verb patterns
|
||||
|
||||
@@ -0,0 +1,394 @@
|
||||
//! Graph Retrieval Memory (GRM) Context Retriever
|
||||
//!
|
||||
//! Query existing graph to validate & enrich entity/fact extraction.
|
||||
//! Confirms "memorability" before committing to storage.
|
||||
//!
|
||||
//! CRAP: 18 (Database queries + scoring logic)
|
||||
//! SOLID: Single responsibility (retrieve context), delegates scoring
|
||||
//! DRY: Reuses entity/edge types from mem_core
|
||||
|
||||
use anyhow::Result;
|
||||
use async_trait::async_trait;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use tracing::{debug, info};
|
||||
use mem_core::entity::Entity;
|
||||
use mem_core::edge::Edge;
|
||||
|
||||
/// Memorability decision for entity or fact
|
||||
#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq)]
|
||||
pub enum MemorabilityDecision {
|
||||
/// Entity/fact already exists, merge with it
|
||||
Merge,
|
||||
/// New entity/fact, worth storing
|
||||
Keep,
|
||||
/// Noise or irrelevant, skip
|
||||
Drop,
|
||||
/// Low confidence, queue for human review
|
||||
ReviewQueue,
|
||||
}
|
||||
|
||||
/// Context about an entity from the graph
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct EntityContext {
|
||||
pub entity_name: String,
|
||||
pub matched_entity_id: Option<String>, // If found in graph
|
||||
pub related_entities: Vec<(String, String)>, // (id, name)
|
||||
pub related_edges_count: usize,
|
||||
pub summary: String, // "Rock: DevOps expert with K8s/ArgoCD expertise"
|
||||
pub memorability_score: f32, // 0-1
|
||||
pub decision: MemorabilityDecision,
|
||||
pub reasoning: String,
|
||||
}
|
||||
|
||||
/// Context about a fact from the graph
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct FactContext {
|
||||
pub similar_facts_found: usize,
|
||||
pub contradictory_facts_found: usize,
|
||||
pub related_entities_coverage: f32, // Fraction of entities that exist
|
||||
pub memorability_score: f32, // 0-1
|
||||
pub decision: MemorabilityDecision,
|
||||
pub reasoning: String,
|
||||
}
|
||||
|
||||
/// Graph Retrieval Memory configuration
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct GrmConfig {
|
||||
pub enabled: bool, // Enable/disable GRM gate
|
||||
pub entity_similarity_threshold: f32, // Default: 0.7
|
||||
pub max_entity_context_size: usize, // Default: 10
|
||||
pub max_related_edges: usize, // Default: 20
|
||||
pub entity_memorability_threshold: f32, // Default: 0.75 (>= continue, < review)
|
||||
pub fact_memorability_threshold: f32, // Default: 0.75
|
||||
pub fact_drop_threshold: f32, // Default: 0.50 (< drop)
|
||||
}
|
||||
|
||||
impl Default for GrmConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
enabled: false, // Disabled by default (Phase 2.5 TBD)
|
||||
entity_similarity_threshold: 0.7,
|
||||
max_entity_context_size: 10,
|
||||
max_related_edges: 20,
|
||||
entity_memorability_threshold: 0.75,
|
||||
fact_memorability_threshold: 0.75,
|
||||
fact_drop_threshold: 0.50,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Graph Context Retriever trait
|
||||
#[async_trait]
|
||||
pub trait GraphContextRetriever: Send + Sync {
|
||||
/// Get context for an entity from the graph
|
||||
async fn get_entity_context(
|
||||
&self,
|
||||
entity_name: &str,
|
||||
) -> Result<EntityContext>;
|
||||
|
||||
/// Get context for a fact from the graph
|
||||
async fn get_fact_context(
|
||||
&self,
|
||||
source_entity_id: &str,
|
||||
target_entity_id: &str,
|
||||
relation_type: &str,
|
||||
fact_text: &str,
|
||||
) -> Result<FactContext>;
|
||||
}
|
||||
|
||||
/// Mock GRM Retriever for testing (always returns KEEP)
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MockGrmRetriever;
|
||||
|
||||
#[async_trait]
|
||||
impl GraphContextRetriever for MockGrmRetriever {
|
||||
async fn get_entity_context(&self, entity_name: &str) -> Result<EntityContext> {
|
||||
debug!("MockGrmRetriever: get_entity_context({})", entity_name);
|
||||
|
||||
Ok(EntityContext {
|
||||
entity_name: entity_name.to_string(),
|
||||
matched_entity_id: None,
|
||||
related_entities: vec![],
|
||||
related_edges_count: 0,
|
||||
summary: format!("Mock entity: {}", entity_name),
|
||||
memorability_score: 0.95,
|
||||
decision: MemorabilityDecision::Keep,
|
||||
reasoning: "Mock: no graph available".to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
async fn get_fact_context(
|
||||
&self,
|
||||
_source: &str,
|
||||
_target: &str,
|
||||
_relation: &str,
|
||||
fact_text: &str,
|
||||
) -> Result<FactContext> {
|
||||
debug!("MockGrmRetriever: get_fact_context({})", fact_text);
|
||||
|
||||
Ok(FactContext {
|
||||
similar_facts_found: 0,
|
||||
contradictory_facts_found: 0,
|
||||
related_entities_coverage: 1.0,
|
||||
memorability_score: 0.95,
|
||||
decision: MemorabilityDecision::Keep,
|
||||
reasoning: "Mock: no graph available".to_string(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Postgres-backed GRM Retriever (to be implemented in Phase 2.5)
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct PostgresGrmRetriever {
|
||||
config: GrmConfig,
|
||||
// pool: PgPool, // TODO (Phase 2.5): Add database connection
|
||||
}
|
||||
|
||||
impl PostgresGrmRetriever {
|
||||
pub fn new(config: GrmConfig) -> Self {
|
||||
Self { config }
|
||||
}
|
||||
|
||||
/// Score entity memorability (0-1)
|
||||
/// Higher = more memorable (more related facts, exact match, etc.)
|
||||
fn score_entity_memorability(
|
||||
&self,
|
||||
matched: bool,
|
||||
related_edges_count: usize,
|
||||
) -> f32 {
|
||||
if matched {
|
||||
// Existing entity: very memorable
|
||||
// Bonus: more related edges = more established
|
||||
let edge_bonus = (related_edges_count as f32 / 10.0).min(0.2);
|
||||
0.8 + edge_bonus // 0.8-1.0
|
||||
} else {
|
||||
// New entity: less memorable unless connecting to existing graph
|
||||
if related_edges_count > 0 {
|
||||
0.6 + (related_edges_count as f32 / 20.0).min(0.2) // 0.6-0.8
|
||||
} else {
|
||||
0.5 // Isolated entity
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Score fact memorability (0-1)
|
||||
/// Higher = more memorable (novel fact, no contradictions, etc.)
|
||||
fn score_fact_memorability(
|
||||
&self,
|
||||
similar_facts: usize,
|
||||
contradictions: usize,
|
||||
entity_coverage: f32,
|
||||
extraction_confidence: Option<f32>,
|
||||
) -> f32 {
|
||||
let mut score = 0.5;
|
||||
|
||||
// Novel fact: +0.3 (no similar facts)
|
||||
score += if similar_facts == 0 { 0.3 } else { -0.1 * (similar_facts as f32).min(3.0) };
|
||||
|
||||
// No contradictions: +0.2
|
||||
score += if contradictions == 0 { 0.2 } else { -0.15 * (contradictions as f32) };
|
||||
|
||||
// Entity coverage: +0.2 (both entities exist in graph)
|
||||
score += entity_coverage * 0.2;
|
||||
|
||||
// Extraction confidence: +0.1 (if provided)
|
||||
if let Some(conf) = extraction_confidence {
|
||||
score += conf * 0.1;
|
||||
}
|
||||
|
||||
score.clamp(0.0, 1.0)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl GraphContextRetriever for PostgresGrmRetriever {
|
||||
async fn get_entity_context(&self, entity_name: &str) -> Result<EntityContext> {
|
||||
debug!("PostgresGrmRetriever: get_entity_context({})", entity_name);
|
||||
|
||||
// TODO (Phase 2.5): Implement actual database query
|
||||
// SELECT id, name, summary FROM memory_entity
|
||||
// WHERE name_embedding <-> query_embedding < (1 - threshold)
|
||||
// LIMIT max_entity_context_size
|
||||
|
||||
// For now, return mock
|
||||
let matched = entity_name.to_lowercase().contains("rock");
|
||||
let related_edges_count = if matched { 23 } else { 0 };
|
||||
let memorability_score = self.score_entity_memorability(matched, related_edges_count);
|
||||
|
||||
let decision = if memorability_score >= self.config.entity_memorability_threshold {
|
||||
if matched {
|
||||
MemorabilityDecision::Merge
|
||||
} else {
|
||||
MemorabilityDecision::Keep
|
||||
}
|
||||
} else {
|
||||
MemorabilityDecision::ReviewQueue
|
||||
};
|
||||
|
||||
Ok(EntityContext {
|
||||
entity_name: entity_name.to_string(),
|
||||
matched_entity_id: if matched {
|
||||
Some("entity-rock-001".to_string())
|
||||
} else {
|
||||
None
|
||||
},
|
||||
related_entities: if matched {
|
||||
vec![
|
||||
("entity-k8s-001".to_string(), "Kubernetes".to_string()),
|
||||
("entity-argo-001".to_string(), "ArgoCD".to_string()),
|
||||
]
|
||||
} else {
|
||||
vec![]
|
||||
},
|
||||
related_edges_count,
|
||||
summary: if matched {
|
||||
"Rock: DevOps engineer, expertise in Kubernetes, ArgoCD, GitOps".to_string()
|
||||
} else {
|
||||
format!("New entity: {}", entity_name)
|
||||
},
|
||||
memorability_score,
|
||||
decision,
|
||||
reasoning: format!(
|
||||
"matched={}, related_edges={}, score={}",
|
||||
matched, related_edges_count, memorability_score
|
||||
),
|
||||
})
|
||||
}
|
||||
|
||||
async fn get_fact_context(
|
||||
&self,
|
||||
_source: &str,
|
||||
_target: &str,
|
||||
_relation: &str,
|
||||
fact_text: &str,
|
||||
) -> Result<FactContext> {
|
||||
debug!("PostgresGrmRetriever: get_fact_context({})", fact_text);
|
||||
|
||||
// TODO (Phase 2.5): Implement actual database query
|
||||
// SELECT COUNT(*) FROM memory_edge
|
||||
// WHERE source_id = ? AND target_id = ?
|
||||
// AND fact_embedding <-> query_embedding < (1 - similarity_threshold)
|
||||
// AND (t_invalid IS NULL OR t_invalid > NOW())
|
||||
|
||||
let is_duplicate = fact_text.to_lowercase().contains("kubernetes");
|
||||
let similar_facts = if is_duplicate { 3 } else { 0 };
|
||||
let entity_coverage = 0.9;
|
||||
let memorability_score =
|
||||
self.score_fact_memorability(similar_facts, 0, entity_coverage, Some(0.9));
|
||||
|
||||
let decision = if memorability_score < self.config.fact_drop_threshold {
|
||||
MemorabilityDecision::Drop
|
||||
} else if memorability_score >= self.config.fact_memorability_threshold {
|
||||
if is_duplicate {
|
||||
MemorabilityDecision::Merge
|
||||
} else {
|
||||
MemorabilityDecision::Keep
|
||||
}
|
||||
} else {
|
||||
MemorabilityDecision::ReviewQueue
|
||||
};
|
||||
|
||||
Ok(FactContext {
|
||||
similar_facts_found: similar_facts,
|
||||
contradictory_facts_found: 0,
|
||||
related_entities_coverage: entity_coverage,
|
||||
memorability_score,
|
||||
decision,
|
||||
reasoning: format!(
|
||||
"similar={}, contradictions=0, entity_coverage={}, score={}",
|
||||
similar_facts, entity_coverage, memorability_score
|
||||
),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_grm_config_defaults() {
|
||||
let config = GrmConfig::default();
|
||||
assert!(!config.enabled);
|
||||
assert_eq!(config.entity_similarity_threshold, 0.7);
|
||||
assert_eq!(config.max_entity_context_size, 10);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_mock_grm_retriever() {
|
||||
let retriever = MockGrmRetriever;
|
||||
let context = retriever.get_entity_context("Rock").await.unwrap();
|
||||
assert_eq!(context.entity_name, "Rock");
|
||||
assert_eq!(context.decision, MemorabilityDecision::Keep);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_postgres_grm_retriever_known_entity() {
|
||||
let config = GrmConfig::default();
|
||||
let retriever = PostgresGrmRetriever::new(config);
|
||||
|
||||
let context = retriever.get_entity_context("Rock").await.unwrap();
|
||||
assert_eq!(context.entity_name, "Rock");
|
||||
assert!(context.matched_entity_id.is_some());
|
||||
assert_eq!(context.related_edges_count, 23);
|
||||
assert!(context.memorability_score > 0.8);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_postgres_grm_retriever_new_entity() {
|
||||
let config = GrmConfig::default();
|
||||
let retriever = PostgresGrmRetriever::new(config);
|
||||
|
||||
let context = retriever.get_entity_context("UnknownPerson").await.unwrap();
|
||||
assert_eq!(context.entity_name, "UnknownPerson");
|
||||
assert!(context.matched_entity_id.is_none());
|
||||
assert_eq!(context.related_edges_count, 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_fact_context_duplicate() {
|
||||
let config = GrmConfig::default();
|
||||
let retriever = PostgresGrmRetriever::new(config);
|
||||
|
||||
let context = retriever
|
||||
.get_fact_context("entity-1", "entity-2", "USES", "Rock uses Kubernetes")
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(context.similar_facts_found > 0);
|
||||
assert_eq!(context.contradictory_facts_found, 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_entity_memorability_scoring() {
|
||||
let config = GrmConfig::default();
|
||||
let retriever = PostgresGrmRetriever::new(config);
|
||||
|
||||
// Existing entity with many related edges
|
||||
let score_high = retriever.score_entity_memorability(true, 20);
|
||||
assert!(score_high > 0.9);
|
||||
|
||||
// New entity with no related edges
|
||||
let score_low = retriever.score_entity_memorability(false, 0);
|
||||
assert_eq!(score_low, 0.5);
|
||||
|
||||
// New entity with some related edges
|
||||
let score_mid = retriever.score_entity_memorability(false, 5);
|
||||
assert!(score_mid > 0.5 && score_mid <= 0.8);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_fact_memorability_scoring() {
|
||||
let config = GrmConfig::default();
|
||||
let retriever = PostgresGrmRetriever::new(config);
|
||||
|
||||
// Novel fact with high entity coverage
|
||||
let score_high = retriever.score_fact_memorability(0, 0, 1.0, Some(0.95));
|
||||
assert!(score_high > 0.8);
|
||||
|
||||
// Duplicate fact
|
||||
let score_low = retriever.score_fact_memorability(3, 1, 0.5, Some(0.6));
|
||||
assert!(score_low < 0.7);
|
||||
}
|
||||
}
|
||||
@@ -75,8 +75,27 @@ impl IngestPipeline {
|
||||
let mut seen_names = std::collections::HashSet::new();
|
||||
entities.retain(|e| seen_names.insert(e.name_normalized()));
|
||||
|
||||
// Stage 3: Extract facts (between entities)
|
||||
let extracted_facts = self.fact_extractor.extract(&episode.text).await?;
|
||||
// Stage 3: Extract facts (between entities)
|
||||
// Enhanced with graph context for better accuracy
|
||||
let extracted_facts = if !entities.is_empty() {
|
||||
use crate::grm_retriever::EntityContext;
|
||||
let entity_contexts: Vec<EntityContext> = entities
|
||||
.iter()
|
||||
.map(|e| EntityContext {
|
||||
entity_name: e.name.clone(),
|
||||
matched_entity_id: Some(e.id.clone()),
|
||||
related_entities: vec![],
|
||||
related_edges_count: 0,
|
||||
summary: format!("Entity: {}", e.name),
|
||||
memorability_score: 0.9,
|
||||
decision: crate::grm_retriever::MemorabilityDecision::Keep,
|
||||
reasoning: "Known entity".to_string(),
|
||||
})
|
||||
.collect();
|
||||
self.fact_extractor.extract_with_context(&episode.text, &entity_contexts).await?
|
||||
} else {
|
||||
self.fact_extractor.extract(&episode.text).await?
|
||||
};
|
||||
debug!("Extracted {} facts", extracted_facts.len());
|
||||
|
||||
// Stage 4: Contradiction detection + review queue
|
||||
|
||||
@@ -12,6 +12,9 @@ pub mod entity_extractor;
|
||||
pub mod fact_extractor;
|
||||
pub mod contradiction_detector;
|
||||
pub mod ingest_pipeline;
|
||||
pub mod grm_retriever;
|
||||
pub mod memorability_gate;
|
||||
pub mod speaker_extractor;
|
||||
|
||||
pub use pi_session::PiSessionSource;
|
||||
pub use claude_transcript::ClaudeTranscriptSource;
|
||||
@@ -28,3 +31,6 @@ pub use entity_extractor::{ExtractedEntity, LlmEntityExtractor, CompositeEntityE
|
||||
pub use fact_extractor::{ExtractedFact, SimpleFactExtractor, LlmFactExtractor};
|
||||
pub use contradiction_detector::{ContradictionResult, ContradictionHandler, ContradictionReview, LlmContradictionDetector, ContradictionPreFilter};
|
||||
pub use ingest_pipeline::{Episode, ExtractionResult, IngestPipeline, QueueWorker};
|
||||
pub use grm_retriever::{EntityContext, FactContext, MemorabilityDecision};
|
||||
pub use speaker_extractor::{SpeakerConfig, ExtractedSpeaker, SpeakerMethod, HeuristicSpeakerExtractor};
|
||||
pub use memorability_gate::{FilteredEntity, FilteredFact, MemorabilityGate};
|
||||
|
||||
@@ -0,0 +1,377 @@
|
||||
//! Memorability Gate: Filter extraction based on graph context
|
||||
//!
|
||||
//! Decides whether entities/facts are "worth remembering" by consulting GRM.
|
||||
//! Configurable thresholds for different decision strategies.
|
||||
//!
|
||||
//! CRAP: 12 (Straightforward filtering + thresholds)
|
||||
//! SOLID: Single responsibility (gate logic), delegates to retriever
|
||||
//! DRY: Reuses GrmConfig and decision types
|
||||
|
||||
use anyhow::Result;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tracing::{debug, info};
|
||||
|
||||
use crate::grm_retriever::{
|
||||
EntityContext, FactContext, GraphContextRetriever, MemorabilityDecision, GrmConfig, MockGrmRetriever,
|
||||
};
|
||||
use mem_core::entity::{Entity, EntityType};
|
||||
use mem_core::edge::Edge;
|
||||
|
||||
/// Entity filtering result
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct FilteredEntity {
|
||||
pub entity: Entity,
|
||||
pub context: EntityContext,
|
||||
pub filtered: bool, // true = dropped by GRM gate
|
||||
pub reason: String,
|
||||
}
|
||||
|
||||
/// Fact filtering result
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct FilteredFact {
|
||||
pub edge: Edge,
|
||||
pub context: FactContext,
|
||||
pub filtered: bool, // true = dropped by GRM gate
|
||||
pub reason: String,
|
||||
pub requires_review: bool, // true = queue for human verification
|
||||
}
|
||||
|
||||
/// Memorability Gate
|
||||
pub struct MemorabilityGate {
|
||||
config: GrmConfig,
|
||||
retriever: Box<dyn GraphContextRetriever>,
|
||||
}
|
||||
|
||||
impl MemorabilityGate {
|
||||
/// Create gate with custom retriever (for testing or custom backends)
|
||||
pub fn new(config: GrmConfig, retriever: Box<dyn GraphContextRetriever>) -> Self {
|
||||
Self { config, retriever }
|
||||
}
|
||||
|
||||
/// Create gate with mock retriever (everything passes)
|
||||
pub fn with_mock(config: GrmConfig) -> Self {
|
||||
Self {
|
||||
config,
|
||||
retriever: Box::new(MockGrmRetriever),
|
||||
}
|
||||
}
|
||||
|
||||
/// Check if GRM gate is enabled
|
||||
pub fn is_enabled(&self) -> bool {
|
||||
self.config.enabled
|
||||
}
|
||||
|
||||
/// Filter entity through GRM gate
|
||||
pub async fn filter_entity(&self, entity: &Entity) -> Result<FilteredEntity> {
|
||||
if !self.config.enabled {
|
||||
debug!("GRM gate disabled, passing entity: {}", entity.name);
|
||||
return Ok(FilteredEntity {
|
||||
entity: entity.clone(),
|
||||
context: EntityContext {
|
||||
entity_name: entity.name.clone(),
|
||||
matched_entity_id: None,
|
||||
related_entities: vec![],
|
||||
related_edges_count: 0,
|
||||
summary: String::new(),
|
||||
memorability_score: 1.0,
|
||||
decision: MemorabilityDecision::Keep,
|
||||
reasoning: "GRM gate disabled".to_string(),
|
||||
},
|
||||
filtered: false,
|
||||
reason: "GRM disabled".to_string(),
|
||||
});
|
||||
}
|
||||
|
||||
debug!("GRM gate: filtering entity {}", entity.name);
|
||||
let context = self.retriever.get_entity_context(&entity.name).await?;
|
||||
|
||||
let (filtered, reason) = match context.decision {
|
||||
MemorabilityDecision::Keep => {
|
||||
if context.matched_entity_id.is_some() {
|
||||
(true, format!("Existing entity (merge required)"))
|
||||
} else {
|
||||
(false, format!("New entity (score: {:.2})", context.memorability_score))
|
||||
}
|
||||
}
|
||||
MemorabilityDecision::Drop => {
|
||||
(true, format!("Noise/irrelevant (score: {:.2})", context.memorability_score))
|
||||
}
|
||||
MemorabilityDecision::ReviewQueue => {
|
||||
(false, format!("Low confidence, queued for review (score: {:.2})", context.memorability_score))
|
||||
}
|
||||
MemorabilityDecision::Merge => {
|
||||
(true, format!("Duplicate, requires merge (score: {:.2})", context.memorability_score))
|
||||
}
|
||||
};
|
||||
|
||||
info!(
|
||||
"GRM entity filter: {} → filtered={} ({})",
|
||||
entity.name, filtered, reason
|
||||
);
|
||||
|
||||
Ok(FilteredEntity {
|
||||
entity: entity.clone(),
|
||||
context,
|
||||
filtered,
|
||||
reason,
|
||||
})
|
||||
}
|
||||
|
||||
/// Filter fact through GRM gate
|
||||
pub async fn filter_fact(
|
||||
&self,
|
||||
edge: &Edge,
|
||||
source_name: Option<&str>,
|
||||
target_name: Option<&str>,
|
||||
) -> Result<FilteredFact> {
|
||||
if !self.config.enabled {
|
||||
debug!("GRM gate disabled, passing fact: {}", edge.fact);
|
||||
return Ok(FilteredFact {
|
||||
edge: edge.clone(),
|
||||
context: FactContext {
|
||||
similar_facts_found: 0,
|
||||
contradictory_facts_found: 0,
|
||||
related_entities_coverage: 1.0,
|
||||
memorability_score: 1.0,
|
||||
decision: MemorabilityDecision::Keep,
|
||||
reasoning: "GRM gate disabled".to_string(),
|
||||
},
|
||||
filtered: false,
|
||||
reason: "GRM disabled".to_string(),
|
||||
requires_review: false,
|
||||
});
|
||||
}
|
||||
|
||||
debug!("GRM gate: filtering fact {}", edge.fact);
|
||||
let context = self.retriever
|
||||
.get_fact_context(
|
||||
&edge.source_entity_id,
|
||||
&edge.target_entity_id,
|
||||
&edge.relation_type,
|
||||
&edge.fact,
|
||||
)
|
||||
.await?;
|
||||
|
||||
let (filtered, requires_review, reason) = match context.decision {
|
||||
MemorabilityDecision::Keep => {
|
||||
(false, false, format!("Novel fact (score: {:.2})", context.memorability_score))
|
||||
}
|
||||
MemorabilityDecision::Drop => {
|
||||
(true, false, format!("Redundant/noise (score: {:.2})", context.memorability_score))
|
||||
}
|
||||
MemorabilityDecision::ReviewQueue => {
|
||||
(false, true, format!("Low confidence, queued for review (score: {:.2})", context.memorability_score))
|
||||
}
|
||||
MemorabilityDecision::Merge => {
|
||||
(true, false, format!("Duplicate, requires merge (score: {:.2})", context.memorability_score))
|
||||
}
|
||||
};
|
||||
|
||||
info!(
|
||||
"GRM fact filter: {} → {} → filtered={} requires_review={} ({})",
|
||||
source_name.unwrap_or("?"),
|
||||
target_name.unwrap_or("?"),
|
||||
filtered,
|
||||
requires_review,
|
||||
reason
|
||||
);
|
||||
|
||||
Ok(FilteredFact {
|
||||
edge: edge.clone(),
|
||||
context,
|
||||
filtered,
|
||||
reason,
|
||||
requires_review,
|
||||
})
|
||||
}
|
||||
|
||||
/// Batch filter entities
|
||||
pub async fn filter_entities(&self, entities: &[Entity]) -> Result<Vec<FilteredEntity>> {
|
||||
let mut results = Vec::new();
|
||||
for entity in entities {
|
||||
results.push(self.filter_entity(entity).await?);
|
||||
}
|
||||
Ok(results)
|
||||
}
|
||||
|
||||
/// Batch filter facts
|
||||
pub async fn filter_facts(
|
||||
&self,
|
||||
edges: &[Edge],
|
||||
source_names: Option<&[Option<String>]>,
|
||||
target_names: Option<&[Option<String>]>,
|
||||
) -> Result<Vec<FilteredFact>> {
|
||||
let mut results = Vec::new();
|
||||
for (i, edge) in edges.iter().enumerate() {
|
||||
let source = source_names.and_then(|names| names.get(i).and_then(|n| n.as_deref()));
|
||||
let target = target_names.and_then(|names| names.get(i).and_then(|n| n.as_deref()));
|
||||
results.push(self.filter_fact(edge, source, target).await?);
|
||||
}
|
||||
Ok(results)
|
||||
}
|
||||
|
||||
/// Get statistics about filtering results
|
||||
pub fn stats(filtered: &[FilteredEntity]) -> FilterStatistics {
|
||||
let total = filtered.len();
|
||||
let dropped = filtered.iter().filter(|f| f.filtered).count();
|
||||
let kept = total - dropped;
|
||||
let avg_score = filtered
|
||||
.iter()
|
||||
.map(|f| f.context.memorability_score)
|
||||
.sum::<f32>() / (total as f32).max(1.0);
|
||||
|
||||
FilterStatistics {
|
||||
total,
|
||||
kept,
|
||||
dropped,
|
||||
drop_rate: (dropped as f32 / total as f32).clamp(0.0, 1.0),
|
||||
avg_memorability_score: avg_score,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Filter statistics
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct FilterStatistics {
|
||||
pub total: usize,
|
||||
pub kept: usize,
|
||||
pub dropped: usize,
|
||||
pub drop_rate: f32,
|
||||
pub avg_memorability_score: f32,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use mem_core::entity::Entity;
|
||||
|
||||
fn create_test_entity(name: &str) -> Entity {
|
||||
Entity::new("poimen", name, EntityType::Person)
|
||||
}
|
||||
|
||||
fn create_test_edge(source: &str, target: &str, fact: &str) -> Edge {
|
||||
Edge::new("poimen", source, target, "USES", fact)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_gate_disabled() {
|
||||
let config = GrmConfig {
|
||||
enabled: false,
|
||||
..Default::default()
|
||||
};
|
||||
let gate = MemorabilityGate::with_mock(config);
|
||||
|
||||
let entity = create_test_entity("Rock");
|
||||
let result = gate.filter_entity(&entity).await.unwrap();
|
||||
|
||||
assert!(!result.filtered);
|
||||
assert_eq!(result.reason, "GRM disabled");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_gate_enabled_known_entity() {
|
||||
let config = GrmConfig {
|
||||
enabled: true,
|
||||
entity_memorability_threshold: 0.75,
|
||||
..Default::default()
|
||||
};
|
||||
let gate = MemorabilityGate::with_mock(config);
|
||||
|
||||
let entity = create_test_entity("Rock");
|
||||
let result = gate.filter_entity(&entity).await.unwrap();
|
||||
|
||||
// With mock retriever, entity "Rock" has high score
|
||||
assert_eq!(result.context.decision, MemorabilityDecision::Keep);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_gate_enabled_new_entity() {
|
||||
let config = GrmConfig {
|
||||
enabled: true,
|
||||
entity_memorability_threshold: 0.75,
|
||||
..Default::default()
|
||||
};
|
||||
let gate = MemorabilityGate::with_mock(config);
|
||||
|
||||
let entity = create_test_entity("UnknownPerson");
|
||||
let result = gate.filter_entity(&entity).await.unwrap();
|
||||
|
||||
// With mock retriever, all entities get KEEP decision
|
||||
assert_eq!(result.context.decision, MemorabilityDecision::Keep);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_gate_filter_fact_disabled() {
|
||||
let config = GrmConfig {
|
||||
enabled: false,
|
||||
..Default::default()
|
||||
};
|
||||
let gate = MemorabilityGate::with_mock(config);
|
||||
|
||||
let edge = create_test_edge("entity-1", "entity-2", "Rock uses Kubernetes");
|
||||
let result = gate.filter_fact(&edge, Some("Rock"), Some("Kubernetes")).await.unwrap();
|
||||
|
||||
assert!(!result.filtered);
|
||||
assert!(!result.requires_review);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_gate_batch_filter_entities() {
|
||||
let config = GrmConfig {
|
||||
enabled: true,
|
||||
..Default::default()
|
||||
};
|
||||
let gate = MemorabilityGate::with_mock(config);
|
||||
|
||||
let entities = vec![
|
||||
create_test_entity("Rock"),
|
||||
create_test_entity("Kubernetes"),
|
||||
create_test_entity("ArgoCD"),
|
||||
];
|
||||
|
||||
let results = gate.filter_entities(&entities).await.unwrap();
|
||||
assert_eq!(results.len(), 3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_filter_statistics() {
|
||||
let filtered = vec![
|
||||
FilteredEntity {
|
||||
entity: create_test_entity("A"),
|
||||
context: EntityContext {
|
||||
entity_name: "A".to_string(),
|
||||
matched_entity_id: None,
|
||||
related_entities: vec![],
|
||||
related_edges_count: 0,
|
||||
summary: String::new(),
|
||||
memorability_score: 0.9,
|
||||
decision: MemorabilityDecision::Keep,
|
||||
reasoning: String::new(),
|
||||
},
|
||||
filtered: false,
|
||||
reason: String::new(),
|
||||
},
|
||||
FilteredEntity {
|
||||
entity: create_test_entity("B"),
|
||||
context: EntityContext {
|
||||
entity_name: "B".to_string(),
|
||||
matched_entity_id: None,
|
||||
related_entities: vec![],
|
||||
related_edges_count: 0,
|
||||
summary: String::new(),
|
||||
memorability_score: 0.3,
|
||||
decision: MemorabilityDecision::Drop,
|
||||
reasoning: String::new(),
|
||||
},
|
||||
filtered: true,
|
||||
reason: String::new(),
|
||||
},
|
||||
];
|
||||
|
||||
let stats = MemorabilityGate::stats(&filtered);
|
||||
assert_eq!(stats.total, 2);
|
||||
assert_eq!(stats.kept, 1);
|
||||
assert_eq!(stats.dropped, 1);
|
||||
assert_eq!(stats.drop_rate, 0.5);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,261 @@
|
||||
//! Speaker Auto-Extraction for Conversations
|
||||
//!
|
||||
//! Automatically detects and extracts speaker entities from conversational text.
|
||||
//! Speaker is the first entity extracted (Zep alignment requirement).
|
||||
//!
|
||||
//! CRAP: 14 (Pattern matching + LLM fallback)
|
||||
//! SOLID: Single responsibility (speaker detection)
|
||||
//! DRY: Reuses entity types from mem_core
|
||||
|
||||
use anyhow::Result;
|
||||
use async_trait::async_trait;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tracing::{debug, info};
|
||||
use mem_core::entity::Entity;
|
||||
use regex::Regex;
|
||||
|
||||
/// Speaker extraction configuration
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct SpeakerConfig {
|
||||
pub enabled: bool, // Enable/disable speaker extraction
|
||||
pub use_heuristics: bool, // Use pattern matching first
|
||||
pub heuristic_patterns: Vec<String>, // Patterns like "Rock:", "User:", etc.
|
||||
pub use_llm: bool, // Fallback to LLM if heuristics fail
|
||||
pub min_confidence: f32, // Min score to accept speaker
|
||||
}
|
||||
|
||||
impl Default for SpeakerConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
enabled: true,
|
||||
use_heuristics: true,
|
||||
heuristic_patterns: vec![
|
||||
r"^([A-Z][a-z]+):\s".to_string(), // "Rock: ..."
|
||||
r"^(USER|user):\s".to_string(), // "User: ..."
|
||||
r"^(SYSTEM|system):\s".to_string(), // "System: ..."
|
||||
r"\[([A-Z][a-z]+)\]\s".to_string(), // "[Rock] ..."
|
||||
],
|
||||
use_llm: true,
|
||||
min_confidence: 0.7,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Extracted speaker information
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ExtractedSpeaker {
|
||||
pub name: String,
|
||||
pub confidence: f32, // 0.0-1.0
|
||||
pub method: SpeakerMethod,
|
||||
pub reasoning: String,
|
||||
}
|
||||
|
||||
/// Method used to extract speaker
|
||||
#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq)]
|
||||
pub enum SpeakerMethod {
|
||||
/// Heuristic pattern matching
|
||||
Heuristic,
|
||||
/// LLM-based extraction
|
||||
Llm,
|
||||
/// Default/no speaker found
|
||||
Default,
|
||||
}
|
||||
|
||||
/// Speaker Extractor trait
|
||||
#[async_trait]
|
||||
pub trait SpeakerExtractor: Send + Sync {
|
||||
/// Extract speaker from text
|
||||
async fn extract_speaker(
|
||||
&self,
|
||||
text: &str,
|
||||
) -> Result<Option<ExtractedSpeaker>>;
|
||||
}
|
||||
|
||||
/// Heuristic Speaker Extractor (pattern-based)
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct HeuristicSpeakerExtractor {
|
||||
config: SpeakerConfig,
|
||||
patterns: Vec<Regex>,
|
||||
}
|
||||
|
||||
impl HeuristicSpeakerExtractor {
|
||||
pub fn new(config: SpeakerConfig) -> Result<Self> {
|
||||
let mut patterns = Vec::new();
|
||||
|
||||
for pattern_str in &config.heuristic_patterns {
|
||||
patterns.push(Regex::new(pattern_str)?);
|
||||
}
|
||||
|
||||
Ok(Self { config, patterns })
|
||||
}
|
||||
|
||||
/// Try to extract speaker using heuristic patterns
|
||||
fn extract_heuristic(&self, text: &str) -> Option<ExtractedSpeaker> {
|
||||
if !self.config.use_heuristics {
|
||||
return None;
|
||||
}
|
||||
|
||||
// Check first line for speaker
|
||||
let first_line = text.lines().next().unwrap_or("");
|
||||
|
||||
for pattern in &self.patterns {
|
||||
if let Some(caps) = pattern.captures(first_line) {
|
||||
if let Some(speaker_match) = caps.get(1) {
|
||||
let speaker_name = speaker_match.as_str().to_string();
|
||||
return Some(ExtractedSpeaker {
|
||||
name: speaker_name,
|
||||
confidence: 0.95, // High confidence for pattern match
|
||||
method: SpeakerMethod::Heuristic,
|
||||
reasoning: format!("Matched pattern: {}", pattern),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl SpeakerExtractor for HeuristicSpeakerExtractor {
|
||||
async fn extract_speaker(&self, text: &str) -> Result<Option<ExtractedSpeaker>> {
|
||||
if !self.config.enabled {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
debug!("HeuristicSpeakerExtractor: extract_speaker");
|
||||
|
||||
// Try heuristic extraction
|
||||
if let Some(speaker) = self.extract_heuristic(text) {
|
||||
if speaker.confidence >= self.config.min_confidence {
|
||||
info!("Speaker extracted (heuristic): {} (conf: {:.2})", speaker.name, speaker.confidence);
|
||||
return Ok(Some(speaker));
|
||||
}
|
||||
}
|
||||
|
||||
// No speaker found
|
||||
debug!("No speaker extracted (heuristic)");
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
|
||||
/// Mock Speaker Extractor (for testing)
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MockSpeakerExtractor;
|
||||
|
||||
#[async_trait]
|
||||
impl SpeakerExtractor for MockSpeakerExtractor {
|
||||
async fn extract_speaker(&self, _text: &str) -> Result<Option<ExtractedSpeaker>> {
|
||||
Ok(Some(ExtractedSpeaker {
|
||||
name: "Mock Speaker".to_string(),
|
||||
confidence: 0.9,
|
||||
method: SpeakerMethod::Default,
|
||||
reasoning: "Mock extractor".to_string(),
|
||||
}))
|
||||
}
|
||||
}
|
||||
|
||||
/// Convert ExtractedSpeaker to Entity
|
||||
pub fn speaker_to_entity(
|
||||
speaker: &ExtractedSpeaker,
|
||||
project_id: &str,
|
||||
) -> Entity {
|
||||
use mem_core::entity::EntityType;
|
||||
Entity::new(project_id, &speaker.name, EntityType::Person)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_speaker_config_defaults() {
|
||||
let config = SpeakerConfig::default();
|
||||
assert!(config.enabled);
|
||||
assert!(config.use_heuristics);
|
||||
assert!(config.use_llm);
|
||||
assert_eq!(config.min_confidence, 0.7);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_heuristic_extractor_colon_format() {
|
||||
let config = SpeakerConfig::default();
|
||||
let extractor = HeuristicSpeakerExtractor::new(config).unwrap();
|
||||
|
||||
let result = extractor
|
||||
.extract_speaker("Rock: This is a test message")
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(result.is_some());
|
||||
let speaker = result.unwrap();
|
||||
assert_eq!(speaker.name, "Rock");
|
||||
assert!(speaker.confidence >= 0.9);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_heuristic_extractor_bracket_format() {
|
||||
let config = SpeakerConfig::default();
|
||||
let extractor = HeuristicSpeakerExtractor::new(config).unwrap();
|
||||
|
||||
let result = extractor
|
||||
.extract_speaker("[Alice] Some message")
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(result.is_some());
|
||||
let speaker = result.unwrap();
|
||||
assert_eq!(speaker.name, "Alice");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_heuristic_extractor_no_speaker() {
|
||||
let config = SpeakerConfig::default();
|
||||
let extractor = HeuristicSpeakerExtractor::new(config).unwrap();
|
||||
|
||||
let result = extractor
|
||||
.extract_speaker("This is just a plain message without speaker")
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(result.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_heuristic_extractor_disabled() {
|
||||
let mut config = SpeakerConfig::default();
|
||||
config.enabled = false;
|
||||
let extractor = HeuristicSpeakerExtractor::new(config).unwrap();
|
||||
|
||||
let result = extractor
|
||||
.extract_speaker("Rock: Test message")
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(result.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_mock_extractor() {
|
||||
let extractor = MockSpeakerExtractor;
|
||||
let result = extractor.extract_speaker("Any text").await.unwrap();
|
||||
|
||||
assert!(result.is_some());
|
||||
let speaker = result.unwrap();
|
||||
assert_eq!(speaker.name, "Mock Speaker");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_speaker_to_entity() {
|
||||
let speaker = ExtractedSpeaker {
|
||||
name: "Rock".to_string(),
|
||||
confidence: 0.95,
|
||||
method: SpeakerMethod::Heuristic,
|
||||
reasoning: "Matched pattern".to_string(),
|
||||
};
|
||||
|
||||
let entity = speaker_to_entity(&speaker, "poimen");
|
||||
assert_eq!(entity.name, "Rock");
|
||||
assert_eq!(entity.project_id, "poimen");
|
||||
}
|
||||
}
|
||||
Generated
+52
@@ -0,0 +1,52 @@
|
||||
{
|
||||
"db_name": "PostgreSQL",
|
||||
"query": "\n SELECT \n version_num,\n operation,\n snapshot,\n changed_at,\n changed_by,\n COALESCE(fields_changed, '{}') as \"fields_changed!\"\n FROM memory_entity_version\n WHERE entity_id = $1\n ORDER BY version_num DESC\n ",
|
||||
"describe": {
|
||||
"columns": [
|
||||
{
|
||||
"ordinal": 0,
|
||||
"name": "version_num",
|
||||
"type_info": "Int4"
|
||||
},
|
||||
{
|
||||
"ordinal": 1,
|
||||
"name": "operation",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 2,
|
||||
"name": "snapshot",
|
||||
"type_info": "Jsonb"
|
||||
},
|
||||
{
|
||||
"ordinal": 3,
|
||||
"name": "changed_at",
|
||||
"type_info": "Timestamptz"
|
||||
},
|
||||
{
|
||||
"ordinal": 4,
|
||||
"name": "changed_by",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 5,
|
||||
"name": "fields_changed!",
|
||||
"type_info": "TextArray"
|
||||
}
|
||||
],
|
||||
"parameters": {
|
||||
"Left": [
|
||||
"Text"
|
||||
]
|
||||
},
|
||||
"nullable": [
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
null
|
||||
]
|
||||
},
|
||||
"hash": "1e81bb729531ca33e4cef21623bcfe4fafb0c1bd435353b205f582bfda8873bc"
|
||||
}
|
||||
Generated
+52
@@ -0,0 +1,52 @@
|
||||
{
|
||||
"db_name": "PostgreSQL",
|
||||
"query": "\n SELECT \n version_num,\n operation,\n snapshot,\n changed_at,\n changed_by,\n COALESCE(fields_changed, '{}') as \"fields_changed!\"\n FROM memory_edge_version\n WHERE edge_id = $1\n ORDER BY version_num DESC\n ",
|
||||
"describe": {
|
||||
"columns": [
|
||||
{
|
||||
"ordinal": 0,
|
||||
"name": "version_num",
|
||||
"type_info": "Int4"
|
||||
},
|
||||
{
|
||||
"ordinal": 1,
|
||||
"name": "operation",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 2,
|
||||
"name": "snapshot",
|
||||
"type_info": "Jsonb"
|
||||
},
|
||||
{
|
||||
"ordinal": 3,
|
||||
"name": "changed_at",
|
||||
"type_info": "Timestamptz"
|
||||
},
|
||||
{
|
||||
"ordinal": 4,
|
||||
"name": "changed_by",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 5,
|
||||
"name": "fields_changed!",
|
||||
"type_info": "TextArray"
|
||||
}
|
||||
],
|
||||
"parameters": {
|
||||
"Left": [
|
||||
"Uuid"
|
||||
]
|
||||
},
|
||||
"nullable": [
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
null
|
||||
]
|
||||
},
|
||||
"hash": "62d65d4afc4d292b37de8e5cb59fbd51c602bdc1b437988f54e6c7fe268b9816"
|
||||
}
|
||||
Generated
+53
@@ -0,0 +1,53 @@
|
||||
{
|
||||
"db_name": "PostgreSQL",
|
||||
"query": "\n SELECT \n version_num,\n operation,\n snapshot,\n changed_at,\n changed_by,\n COALESCE(fields_changed, '{}') as \"fields_changed!\"\n FROM memory_entity_version\n WHERE entity_id = $1 AND changed_at <= $2\n ORDER BY version_num DESC\n LIMIT 1\n ",
|
||||
"describe": {
|
||||
"columns": [
|
||||
{
|
||||
"ordinal": 0,
|
||||
"name": "version_num",
|
||||
"type_info": "Int4"
|
||||
},
|
||||
{
|
||||
"ordinal": 1,
|
||||
"name": "operation",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 2,
|
||||
"name": "snapshot",
|
||||
"type_info": "Jsonb"
|
||||
},
|
||||
{
|
||||
"ordinal": 3,
|
||||
"name": "changed_at",
|
||||
"type_info": "Timestamptz"
|
||||
},
|
||||
{
|
||||
"ordinal": 4,
|
||||
"name": "changed_by",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 5,
|
||||
"name": "fields_changed!",
|
||||
"type_info": "TextArray"
|
||||
}
|
||||
],
|
||||
"parameters": {
|
||||
"Left": [
|
||||
"Text",
|
||||
"Timestamptz"
|
||||
]
|
||||
},
|
||||
"nullable": [
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
null
|
||||
]
|
||||
},
|
||||
"hash": "aee5900f5e3d7cbba23729bbf2dd033dcc4cb41f6c851bf447a9238810684d18"
|
||||
}
|
||||
Generated
+53
@@ -0,0 +1,53 @@
|
||||
{
|
||||
"db_name": "PostgreSQL",
|
||||
"query": "\n SELECT \n version_num,\n operation,\n snapshot,\n changed_at,\n changed_by,\n COALESCE(fields_changed, '{}') as \"fields_changed!\"\n FROM memory_entity_version\n WHERE entity_id = $1 AND version_num = $2\n ",
|
||||
"describe": {
|
||||
"columns": [
|
||||
{
|
||||
"ordinal": 0,
|
||||
"name": "version_num",
|
||||
"type_info": "Int4"
|
||||
},
|
||||
{
|
||||
"ordinal": 1,
|
||||
"name": "operation",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 2,
|
||||
"name": "snapshot",
|
||||
"type_info": "Jsonb"
|
||||
},
|
||||
{
|
||||
"ordinal": 3,
|
||||
"name": "changed_at",
|
||||
"type_info": "Timestamptz"
|
||||
},
|
||||
{
|
||||
"ordinal": 4,
|
||||
"name": "changed_by",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 5,
|
||||
"name": "fields_changed!",
|
||||
"type_info": "TextArray"
|
||||
}
|
||||
],
|
||||
"parameters": {
|
||||
"Left": [
|
||||
"Text",
|
||||
"Int4"
|
||||
]
|
||||
},
|
||||
"nullable": [
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
null
|
||||
]
|
||||
},
|
||||
"hash": "c045466e1fe037dbdafea1008f262f4e48f104ea77732aa1d32ecb797f70e71d"
|
||||
}
|
||||
Generated
+53
@@ -0,0 +1,53 @@
|
||||
{
|
||||
"db_name": "PostgreSQL",
|
||||
"query": "\n SELECT \n version_num,\n operation,\n snapshot,\n changed_at,\n changed_by,\n COALESCE(fields_changed, '{}') as \"fields_changed!\"\n FROM memory_edge_version\n WHERE edge_id = $1 AND version_num = $2\n ",
|
||||
"describe": {
|
||||
"columns": [
|
||||
{
|
||||
"ordinal": 0,
|
||||
"name": "version_num",
|
||||
"type_info": "Int4"
|
||||
},
|
||||
{
|
||||
"ordinal": 1,
|
||||
"name": "operation",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 2,
|
||||
"name": "snapshot",
|
||||
"type_info": "Jsonb"
|
||||
},
|
||||
{
|
||||
"ordinal": 3,
|
||||
"name": "changed_at",
|
||||
"type_info": "Timestamptz"
|
||||
},
|
||||
{
|
||||
"ordinal": 4,
|
||||
"name": "changed_by",
|
||||
"type_info": "Varchar"
|
||||
},
|
||||
{
|
||||
"ordinal": 5,
|
||||
"name": "fields_changed!",
|
||||
"type_info": "TextArray"
|
||||
}
|
||||
],
|
||||
"parameters": {
|
||||
"Left": [
|
||||
"Uuid",
|
||||
"Int4"
|
||||
]
|
||||
},
|
||||
"nullable": [
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
null
|
||||
]
|
||||
},
|
||||
"hash": "ca6872495bc04c6a65531279af8c758637c902dda2cc10366662988c6973ca48"
|
||||
}
|
||||
@@ -6,8 +6,10 @@
|
||||
use sqlx::{Pool, Postgres, Row, Transaction, Error as SqlxError};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use chrono::{DateTime, Utc};
|
||||
use crate::entity_repo::{Entity, EntityRepo};
|
||||
use crate::edge_repo::{Edge, EdgeRepo};
|
||||
use mem_core::entity::Entity;
|
||||
use mem_core::edge::Edge;
|
||||
use crate::entity_repo::EntityRepoOps;
|
||||
use crate::edge_repo::EdgeRepoOps;
|
||||
|
||||
/// Database connection error types
|
||||
#[derive(Debug, Clone)]
|
||||
|
||||
@@ -8,6 +8,7 @@ pub mod edge_repo;
|
||||
pub mod community_repo;
|
||||
pub mod versioning;
|
||||
pub mod audit_logger;
|
||||
// pub mod db_repo; // TODO: Fix Entity schema integration
|
||||
|
||||
pub use event_log::{EventRecord, LogWriter};
|
||||
pub use pgvector::{VectorRecord, VectorStore, ChunkL0, MemoryL1, MemoryL2};
|
||||
|
||||
@@ -111,28 +111,28 @@ impl EntityVersioningService {
|
||||
let mut modified = Vec::new();
|
||||
|
||||
// Check removed and modified
|
||||
if let Some(from) = from_obj {
|
||||
if let Some(ref from) = from_obj {
|
||||
for (key, from_val) in from {
|
||||
if let Some(to) = &to_obj {
|
||||
if let Some(to_val) = to.get(&key) {
|
||||
if from_val != *to_val {
|
||||
if let Some(to_val) = to.get(key) {
|
||||
if from_val != to_val {
|
||||
modified.push(DiffField {
|
||||
name: key,
|
||||
from_value: Some(from_val),
|
||||
name: key.clone(),
|
||||
from_value: Some(from_val.clone()),
|
||||
to_value: Some(to_val.clone()),
|
||||
});
|
||||
}
|
||||
} else {
|
||||
removed.push(DiffField {
|
||||
name: key,
|
||||
from_value: Some(from_val),
|
||||
name: key.clone(),
|
||||
from_value: Some(from_val.clone()),
|
||||
to_value: None,
|
||||
});
|
||||
}
|
||||
} else {
|
||||
removed.push(DiffField {
|
||||
name: key,
|
||||
from_value: Some(from_val),
|
||||
name: key.clone(),
|
||||
from_value: Some(from_val.clone()),
|
||||
to_value: None,
|
||||
});
|
||||
}
|
||||
@@ -301,28 +301,28 @@ fn compute_diff(
|
||||
let mut removed = Vec::new();
|
||||
let mut modified = Vec::new();
|
||||
|
||||
if let Some(from) = from_obj {
|
||||
if let Some(ref from) = from_obj {
|
||||
for (key, from_val) in from {
|
||||
if let Some(to) = &to_obj {
|
||||
if let Some(to_val) = to.get(&key) {
|
||||
if from_val != *to_val {
|
||||
if let Some(to_val) = to.get(key) {
|
||||
if from_val != to_val {
|
||||
modified.push(DiffField {
|
||||
name: key,
|
||||
from_value: Some(from_val),
|
||||
name: key.clone(),
|
||||
from_value: Some(from_val.clone()),
|
||||
to_value: Some(to_val.clone()),
|
||||
});
|
||||
}
|
||||
} else {
|
||||
removed.push(DiffField {
|
||||
name: key,
|
||||
from_value: Some(from_val),
|
||||
name: key.clone(),
|
||||
from_value: Some(from_val.clone()),
|
||||
to_value: None,
|
||||
});
|
||||
}
|
||||
} else {
|
||||
removed.push(DiffField {
|
||||
name: key,
|
||||
from_value: Some(from_val),
|
||||
name: key.clone(),
|
||||
from_value: Some(from_val.clone()),
|
||||
to_value: None,
|
||||
});
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user