Compare commits
4
Commits
2307e5379c
..
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
fb61de6b47 | ||
|
|
6b18d81421 | ||
|
|
b15072e12d | ||
|
|
1e5c3d1433 |
@@ -35,6 +35,9 @@ jobs:
|
|||||||
- name: Cargo clippy
|
- name: Cargo clippy
|
||||||
run: cargo clippy --all --all-targets -- -D warnings 2>&1 | tail -50 || true
|
run: cargo clippy --all --all-targets -- -D warnings 2>&1 | tail -50 || true
|
||||||
|
|
||||||
|
- name: Clean build artifacts before Docker
|
||||||
|
run: cargo clean
|
||||||
|
|
||||||
- name: Get short SHA
|
- name: Get short SHA
|
||||||
id: sha
|
id: sha
|
||||||
run: echo "short_sha=$(git rev-parse --short HEAD)" >> $GITHUB_OUTPUT
|
run: echo "short_sha=$(git rev-parse --short HEAD)" >> $GITHUB_OUTPUT
|
||||||
@@ -47,19 +50,13 @@ jobs:
|
|||||||
REGISTRY_USER: ${{ secrets.FORGEJO_REGISTRY_USER }}
|
REGISTRY_USER: ${{ secrets.FORGEJO_REGISTRY_USER }}
|
||||||
REGISTRY_TOKEN: ${{ secrets.FORGEJO_REGISTRY_TOKEN }}
|
REGISTRY_TOKEN: ${{ secrets.FORGEJO_REGISTRY_TOKEN }}
|
||||||
|
|
||||||
- name: Build Docker image
|
- name: Build and push Docker image (SHA tag only)
|
||||||
run: |
|
run: |
|
||||||
docker build --no-cache --progress=plain \
|
docker build --no-cache --progress=plain \
|
||||||
-t "${IMAGE}:${{ steps.sha.outputs.short_sha }}" \
|
-t "${IMAGE}:${{ steps.sha.outputs.short_sha }}" \
|
||||||
-t "${IMAGE}:latest" \
|
|
||||||
-f Dockerfile .
|
-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}:${{ steps.sha.outputs.short_sha }}"
|
||||||
docker push "${IMAGE}:latest"
|
echo "Pushed: ${IMAGE}:${{ steps.sha.outputs.short_sha }}"
|
||||||
echo "✓ Pushed: ${IMAGE}:${{ steps.sha.outputs.short_sha }}"
|
|
||||||
|
|
||||||
- name: Prune unused images
|
- name: Prune unused images
|
||||||
run: docker image prune -a --force 2>&1 | tail -3 || true
|
run: docker image prune -a --force 2>&1 | tail -3 || true
|
||||||
|
|||||||
@@ -0,0 +1,44 @@
|
|||||||
|
name: Deploy
|
||||||
|
|
||||||
|
on:
|
||||||
|
push:
|
||||||
|
branches: [main]
|
||||||
|
workflow_dispatch:
|
||||||
|
|
||||||
|
env:
|
||||||
|
REGISTRY: forgejo.riotpiao.com
|
||||||
|
IMAGE: forgejo.riotpiao.com/riotpiao-poimen/poimen-memory
|
||||||
|
DOCKER_HOST: tcp://localhost:2375
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
deploy:
|
||||||
|
name: Tag & Push Latest
|
||||||
|
runs-on: rust
|
||||||
|
steps:
|
||||||
|
- name: Install Docker
|
||||||
|
run: apt-get update && apt-get install -y docker.io
|
||||||
|
|
||||||
|
- name: Checkout code
|
||||||
|
uses: actions/checkout@v4
|
||||||
|
|
||||||
|
- name: Get short SHA
|
||||||
|
id: sha
|
||||||
|
run: echo "short_sha=$(git rev-parse --short HEAD)" >> $GITHUB_OUTPUT
|
||||||
|
|
||||||
|
- name: Registry login
|
||||||
|
run: |
|
||||||
|
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: Pull SHA image and tag as latest
|
||||||
|
run: |
|
||||||
|
docker pull "${IMAGE}:${{ steps.sha.outputs.short_sha }}" && \
|
||||||
|
docker tag "${IMAGE}:${{ steps.sha.outputs.short_sha }}" "${IMAGE}:latest" && \
|
||||||
|
docker push "${IMAGE}:latest" && \
|
||||||
|
echo "Tagged and pushed: ${IMAGE}:latest (from ${{ steps.sha.outputs.short_sha }})"
|
||||||
|
|
||||||
|
- name: Prune images
|
||||||
|
run: docker image prune -a --force 2>&1 | tail -3 || true
|
||||||
@@ -0,0 +1,76 @@
|
|||||||
|
name: DB Migration
|
||||||
|
|
||||||
|
on:
|
||||||
|
push:
|
||||||
|
branches: [main]
|
||||||
|
paths:
|
||||||
|
- 'crates/mem-store/migrations/**'
|
||||||
|
workflow_dispatch:
|
||||||
|
|
||||||
|
env:
|
||||||
|
DB_HOST: memory-db-rw.poimen.svc.cluster.local
|
||||||
|
DB_PORT: "5432"
|
||||||
|
DB_NAME: memory
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
migrate:
|
||||||
|
name: Run Migrations
|
||||||
|
runs-on: rust
|
||||||
|
steps:
|
||||||
|
- name: Install psql
|
||||||
|
run: apt-get update && apt-get install -y postgresql-client
|
||||||
|
|
||||||
|
- name: Checkout code
|
||||||
|
uses: actions/checkout@v4
|
||||||
|
|
||||||
|
- name: Fetch previous migrations state
|
||||||
|
run: |
|
||||||
|
git fetch origin main --depth=2
|
||||||
|
# List changed migration files
|
||||||
|
CHANGED=$(git diff --name-only HEAD~1 HEAD -- crates/mem-store/migrations/ || echo "")
|
||||||
|
echo "Changed migrations: $CHANGED"
|
||||||
|
echo "CHANGED_MIGRATIONS=$CHANGED" >> $GITHUB_ENV
|
||||||
|
|
||||||
|
- name: Run migrations
|
||||||
|
if: env.CHANGED_MIGRATIONS != ''
|
||||||
|
run: |
|
||||||
|
export PGPASSWORD="${DB_PASSWORD}"
|
||||||
|
|
||||||
|
echo "=== Running changed migrations ==="
|
||||||
|
for f in $CHANGED_MIGRATIONS; do
|
||||||
|
if [ -f "$f" ]; then
|
||||||
|
echo "--- Applying: $f ---"
|
||||||
|
psql -h "$DB_HOST" -p "$DB_PORT" -U "$DB_USER" -d "$DB_NAME" -f "$f" 2>&1
|
||||||
|
if [ $? -ne 0 ]; then
|
||||||
|
echo "ERROR: Migration $f failed!"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
echo "--- OK: $f ---"
|
||||||
|
fi
|
||||||
|
done
|
||||||
|
|
||||||
|
echo "=== Verify schema ==="
|
||||||
|
psql -h "$DB_HOST" -p "$DB_PORT" -U "$DB_USER" -d "$DB_NAME" -c "\dt memory*"
|
||||||
|
env:
|
||||||
|
DB_USER: ${{ secrets.DB_USER }}
|
||||||
|
DB_PASSWORD: ${{ secrets.DB_PASSWORD }}
|
||||||
|
|
||||||
|
- name: Run all migrations (manual trigger)
|
||||||
|
if: github.event_name == 'workflow_dispatch'
|
||||||
|
run: |
|
||||||
|
export PGPASSWORD="${DB_PASSWORD}"
|
||||||
|
|
||||||
|
echo "=== Running all migrations in order ==="
|
||||||
|
for f in $(ls crates/mem-store/migrations/*.sql | sort); do
|
||||||
|
echo "--- Applying: $f ---"
|
||||||
|
psql -h "$DB_HOST" -p "$DB_PORT" -U "$DB_USER" -d "$DB_NAME" -f "$f" 2>&1 || true
|
||||||
|
echo "--- Done: $f ---"
|
||||||
|
done
|
||||||
|
|
||||||
|
echo "=== Final schema ==="
|
||||||
|
psql -h "$DB_HOST" -p "$DB_PORT" -U "$DB_USER" -d "$DB_NAME" -c "\dt memory*"
|
||||||
|
psql -h "$DB_HOST" -p "$DB_PORT" -U "$DB_USER" -d "$DB_NAME" -c "\d memory_entity"
|
||||||
|
psql -h "$DB_HOST" -p "$DB_PORT" -U "$DB_USER" -d "$DB_NAME" -c "\d memory_edge"
|
||||||
|
env:
|
||||||
|
DB_USER: ${{ secrets.DB_USER }}
|
||||||
|
DB_PASSWORD: ${{ secrets.DB_PASSWORD }}
|
||||||
@@ -0,0 +1,136 @@
|
|||||||
|
# Poimen Memory System
|
||||||
|
|
||||||
|
## Project Status
|
||||||
|
|
||||||
|
**Architecture**: Temporal Knowledge Graph for Agent Memory (Zep paper alignment — arXiv:2501.13956)
|
||||||
|
|
||||||
|
**Current**: Ingest pipeline with LLM entity + fact extraction working E2E. Deployed to K8s.
|
||||||
|
|
||||||
|
### What Works
|
||||||
|
- ✅ HTTP server (actix-web) with 15+ endpoints
|
||||||
|
- ✅ LLM entity extraction (LlmEntityExtractor) — extracts person/tool/concept/org entities
|
||||||
|
- ✅ LLM fact extraction (LlmFactExtractor) — extracts relationships between entities
|
||||||
|
- ✅ Reasoning model support — strips `<think>` tags, markdown fences
|
||||||
|
- ✅ Ollama + vLLM + OpenAI-compatible API support
|
||||||
|
- ✅ Entity persistence to pgvector (memory_entity table)
|
||||||
|
- ✅ Edge persistence (memory_edge table with temporal fields)
|
||||||
|
- ✅ Graph query endpoints (entities, edges, BFS traversal)
|
||||||
|
- ✅ Visualization (React Flow JSON, force-directed layout, SSE streaming)
|
||||||
|
- ✅ JWT auth (Authentik OIDC) with RBAC
|
||||||
|
- ✅ K8s deployment (CNPG postgres, ConfigMap, SOPS secrets)
|
||||||
|
- ✅ CI: PR builds push :SHA tag, main merges retag :latest
|
||||||
|
- ✅ 781 tests passing
|
||||||
|
|
||||||
|
### Deployment
|
||||||
|
- **Namespace**: `poimen`
|
||||||
|
- **Image**: `forgejo.riotpiao.com/riotpiao-poimen/poimen-memory:latest`
|
||||||
|
- **DB**: CNPG cluster `memory-db` (pgvector)
|
||||||
|
- **LLM**: `reasoning-predictor.llm-serving.svc.cluster.local` (ornith:35b / qwen2.5:3b)
|
||||||
|
- **Auth**: Authentik OIDC (`MEM_AUTH_MODE=none` for dev)
|
||||||
|
- **Registry**: Forgejo container registry (FORGEJO_REGISTRY_USER/TOKEN secrets)
|
||||||
|
|
||||||
|
### Key Env Vars
|
||||||
|
```
|
||||||
|
DATABASE_URL postgresql://...
|
||||||
|
MEM_AUTH_MODE none|jwt|apikey
|
||||||
|
LLM_ENDPOINT http://localhost:11434/v1/chat/completions (Ollama)
|
||||||
|
LLM_MODEL qwen2.5:3b | ornith:35b | reasoning
|
||||||
|
LLM_API_KEY (for authenticated LLM APIs)
|
||||||
|
MEM_API_KEY (server API key, fallback "test-key")
|
||||||
|
OPENSEARCH_HOSTS (optional, hybrid search)
|
||||||
|
GATEWAY_URL (optional, external queue)
|
||||||
|
```
|
||||||
|
|
||||||
|
## Rules
|
||||||
|
|
||||||
|
1. **No progress markdown files.** Track via Forgejo issues + PRs only.
|
||||||
|
2. **Obsidian vault repo**: `ssh://[email protected]:2222/rock/poimen-obesdient-memory.git`
|
||||||
|
3. **Secrets via KSOPS**: Age-based SOPS encryption. Never commit plaintext.
|
||||||
|
4. **Tea CLI**: `poimen` login has API token `1f717a00134f17c9d2d656c620b955e03ea41276`
|
||||||
|
|
||||||
|
## Architecture (Zep Paper §2)
|
||||||
|
|
||||||
|
### Three-Tier Knowledge Graph
|
||||||
|
```
|
||||||
|
Episode Subgraph (raw messages)
|
||||||
|
→ Entity Subgraph (extracted entities + facts/edges)
|
||||||
|
→ Community Subgraph (clusters, planned Phase 4)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Ingest Pipeline (4 stages)
|
||||||
|
1. **Entity extraction** — LLM extracts named entities with type + summary
|
||||||
|
2. **Deduplication** — HashSet on normalized name
|
||||||
|
3. **Fact extraction** — LLM extracts relationships between entity pairs
|
||||||
|
4. **Contradiction detection** — pre-filter + review queue
|
||||||
|
|
||||||
|
### Retrieval (3 methods, §3)
|
||||||
|
- Cosine semantic similarity (pgvector HNSW)
|
||||||
|
- BM25 full-text (OpenSearch, optional)
|
||||||
|
- BFS graph traversal (depth 1-3)
|
||||||
|
|
||||||
|
### Extractors
|
||||||
|
- `LlmEntityExtractor`: calls LLM_ENDPOINT, parses JSON, handles reasoning models
|
||||||
|
- `LlmFactExtractor`: takes entity list + text, extracts edges between known entities
|
||||||
|
- `WikiLinkFallbackExtractor`: pattern-matches `[[wiki links]]` (no LLM)
|
||||||
|
- `SimpleFactExtractor`: verb pattern matching (no LLM)
|
||||||
|
- Selection: LLM extractors when `LLM_ENDPOINT` set, else fallbacks
|
||||||
|
|
||||||
|
### LLM Response Cleaning
|
||||||
|
`clean_llm_response()` handles:
|
||||||
|
- `<think>...</think>` blocks (reasoning models)
|
||||||
|
- Markdown code fences (```json ... ```)
|
||||||
|
- Array responses (wrap in `{"entities": [...]}`)
|
||||||
|
- Extract first JSON object from mixed text
|
||||||
|
|
||||||
|
## Crate Structure
|
||||||
|
|
||||||
|
```
|
||||||
|
crates/
|
||||||
|
mem-core/ — Entity, Edge, domain types (174 tests)
|
||||||
|
mem-store/ — DB repos, schema, vector store
|
||||||
|
mem-ingest/ — Entity/fact extraction, contradiction detection (87 tests)
|
||||||
|
mem-llm/ — Embeddings, chat, rerank clients
|
||||||
|
mem-cli/ — HTTP server, handlers, query, ingest worker (496 tests)
|
||||||
|
```
|
||||||
|
|
||||||
|
## API Endpoints
|
||||||
|
|
||||||
|
```
|
||||||
|
GET /health
|
||||||
|
POST /memory/ingest — Queue ingest job
|
||||||
|
GET /memory/ingest/{id} — Check job status
|
||||||
|
GET /memory/query?project=&question= — Graph query
|
||||||
|
POST /memory/query — Unified query
|
||||||
|
POST /memory/context — Three-tier retrieval
|
||||||
|
POST /memory/learn — Direct learn
|
||||||
|
POST /memory/visualize — React Flow JSON
|
||||||
|
POST /memory/visualize/stream — SSE streaming
|
||||||
|
POST /memory/compact — Trigger compaction
|
||||||
|
GET /memory/projects — List projects
|
||||||
|
GET /memory/skills — List skills
|
||||||
|
GET /memory/vault — Browse vault
|
||||||
|
POST /memory/synthesis/* — Entity linking, alias detection
|
||||||
|
```
|
||||||
|
|
||||||
|
## Current PRs / Branches
|
||||||
|
|
||||||
|
- **PR #48** `feat/memory-ingest-retrieval` — LLM entity + fact extraction, deployment fixes
|
||||||
|
- **PR #47** merged — Agent entity types (Phase 3.1)
|
||||||
|
- **PR #46** merged — Integration test fixes, CI
|
||||||
|
|
||||||
|
## Next Steps
|
||||||
|
|
||||||
|
1. Merge PR #48 → new image with LLM extraction
|
||||||
|
2. Query retrieval E2E — verify entities/edges returned in query results
|
||||||
|
3. Visualization E2E — test /memory/visualize with extracted graph
|
||||||
|
4. Restore 198 deleted tests from PR #46
|
||||||
|
5. Community detection (Phase 4, Zep §2.3)
|
||||||
|
6. Temporal edge invalidation (Zep §2.2.3)
|
||||||
|
7. Reranker (cross-encoder, RRF, episode-mentions — Zep §3.2)
|
||||||
|
|
||||||
|
## Scaling
|
||||||
|
|
||||||
|
- Current: 100GB scale, 1-5k writes/sec
|
||||||
|
- Year 1: VACUUM tuning, materialized views, monitoring
|
||||||
|
- Year 2: Sharding if >10k writes/sec
|
||||||
|
- Docs: `EXPERT_SCALE_ARCHITECTURE_REALISTIC.md`
|
||||||
+3
-1
@@ -10,7 +10,9 @@ COPY . .
|
|||||||
|
|
||||||
# Build the mem binary (offline sqlx - uses .sqlx/ cache)
|
# Build the mem binary (offline sqlx - uses .sqlx/ cache)
|
||||||
ENV SQLX_OFFLINE=true
|
ENV SQLX_OFFLINE=true
|
||||||
RUN cargo build --release -p mem-cli
|
RUN cargo build --release -p mem-cli && \
|
||||||
|
strip target/release/mem && \
|
||||||
|
rm -rf target/release/deps target/release/build target/release/incremental target/release/.fingerprint
|
||||||
|
|
||||||
# Stage 2: Runtime
|
# Stage 2: Runtime
|
||||||
FROM debian:bookworm-slim
|
FROM debian:bookworm-slim
|
||||||
|
|||||||
@@ -126,246 +126,3 @@ impl Default for MetricsCollector {
|
|||||||
// - Only record_request() needs exclusive write lock
|
// - Only record_request() needs exclusive write lock
|
||||||
// - Performance improvement for high-read scenarios
|
// - Performance improvement for high-read scenarios
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_agent_metrics_default() {
|
|
||||||
let m = AgentMetrics::default();
|
|
||||||
assert_eq!(m.requests_total, 0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_agent_metrics_creation() {
|
|
||||||
let m = AgentMetrics {
|
|
||||||
agent_id: "a1".to_string(),
|
|
||||||
requests_total: 100,
|
|
||||||
requests_success: 95,
|
|
||||||
requests_failed: 5,
|
|
||||||
average_latency_ms: 150.0,
|
|
||||||
p95_latency_ms: 300.0,
|
|
||||||
p99_latency_ms: 450.0,
|
|
||||||
capabilities_used: HashMap::new(),
|
|
||||||
last_updated: "2025-01-30T10:00:00Z".to_string(),
|
|
||||||
};
|
|
||||||
assert_eq!(m.requests_total, 100);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_metrics_collector_creation() {
|
|
||||||
let collector = MetricsCollector::new();
|
|
||||||
assert!(collector.get_metrics("unknown").is_none());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_metrics_collector_concurrent_reads() {
|
|
||||||
let collector = std::sync::Arc::new(MetricsCollector::new());
|
|
||||||
collector.record_request("agent1", true, 100.0, None);
|
|
||||||
|
|
||||||
let mut handles = vec![];
|
|
||||||
for _ in 0..5 {
|
|
||||||
let c = collector.clone();
|
|
||||||
let handle = std::thread::spawn(move || {
|
|
||||||
c.get_metrics("agent1")
|
|
||||||
});
|
|
||||||
handles.push(handle);
|
|
||||||
}
|
|
||||||
|
|
||||||
for handle in handles {
|
|
||||||
assert!(handle.join().unwrap().is_some());
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_metrics_collector_record_success() {
|
|
||||||
let collector = MetricsCollector::new();
|
|
||||||
collector.record_request("agent1", true, 100.0, Some("synthesis"));
|
|
||||||
|
|
||||||
let metrics = collector.get_metrics("agent1");
|
|
||||||
assert!(metrics.is_some());
|
|
||||||
let m = metrics.unwrap();
|
|
||||||
assert_eq!(m.requests_total, 1);
|
|
||||||
assert_eq!(m.requests_success, 1);
|
|
||||||
assert_eq!(m.requests_failed, 0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_metrics_success_rate_calc() {
|
|
||||||
let collector = MetricsCollector::new();
|
|
||||||
for _ in 0..9 {
|
|
||||||
collector.record_request("agent1", true, 100.0, None);
|
|
||||||
}
|
|
||||||
collector.record_request("agent1", false, 50.0, None);
|
|
||||||
|
|
||||||
let m = collector.get_metrics("agent1").unwrap();
|
|
||||||
let success_rate = m.requests_success as f32 / m.requests_total as f32;
|
|
||||||
assert!((success_rate - 0.9).abs() < 0.01);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_metrics_collector_record_failure() {
|
|
||||||
let collector = MetricsCollector::new();
|
|
||||||
collector.record_request("agent1", false, 50.0, None);
|
|
||||||
|
|
||||||
let metrics = collector.get_metrics("agent1");
|
|
||||||
let m = metrics.unwrap();
|
|
||||||
assert_eq!(m.requests_failed, 1);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_metrics_no_contention() {
|
|
||||||
let collector = std::sync::Arc::new(MetricsCollector::new());
|
|
||||||
let mut handles = vec![];
|
|
||||||
|
|
||||||
for i in 0..5 {
|
|
||||||
let c = collector.clone();
|
|
||||||
let h1 = std::thread::spawn(move || {
|
|
||||||
c.record_request(&format!("agent{}", i), true, 100.0, None);
|
|
||||||
});
|
|
||||||
handles.push(h1);
|
|
||||||
|
|
||||||
let c = collector.clone();
|
|
||||||
let h2 = std::thread::spawn(move || {
|
|
||||||
c.get_metrics(&format!("agent{}", i))
|
|
||||||
});
|
|
||||||
handles.push(h2);
|
|
||||||
}
|
|
||||||
|
|
||||||
for h in handles {
|
|
||||||
h.join().unwrap();
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_metrics_collector_multiple_records() {
|
|
||||||
let collector = MetricsCollector::new();
|
|
||||||
collector.record_request("agent1", true, 100.0, None);
|
|
||||||
collector.record_request("agent1", true, 150.0, None);
|
|
||||||
collector.record_request("agent1", false, 50.0, None);
|
|
||||||
|
|
||||||
let metrics = collector.get_metrics("agent1");
|
|
||||||
let m = metrics.unwrap();
|
|
||||||
assert_eq!(m.requests_total, 3);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_metrics_fail_count() {
|
|
||||||
let collector = MetricsCollector::new();
|
|
||||||
collector.record_request("agent1", false, 100.0, None);
|
|
||||||
collector.record_request("agent1", false, 120.0, None);
|
|
||||||
|
|
||||||
let metrics = collector.get_metrics("agent1").unwrap();
|
|
||||||
assert_eq!(metrics.requests_failed, 2);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_metrics_collector_capability_tracking() {
|
|
||||||
let collector = MetricsCollector::new();
|
|
||||||
collector.record_request("agent1", true, 100.0, Some("linking"));
|
|
||||||
collector.record_request("agent1", true, 120.0, Some("linking"));
|
|
||||||
collector.record_request("agent1", true, 110.0, Some("inference"));
|
|
||||||
|
|
||||||
let metrics = collector.get_metrics("agent1");
|
|
||||||
let m = metrics.unwrap();
|
|
||||||
assert_eq!(m.capabilities_used.get("linking"), Some(&2));
|
|
||||||
assert_eq!(m.capabilities_used.get("inference"), Some(&1));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_metrics_thread_safety() {
|
|
||||||
let collector = std::sync::Arc::new(MetricsCollector::new());
|
|
||||||
let mut handles = vec![];
|
|
||||||
|
|
||||||
for i in 0..10 {
|
|
||||||
let c = collector.clone();
|
|
||||||
let handle = std::thread::spawn(move || {
|
|
||||||
c.record_request(&format!("agent{}", i), true, 100.0, None);
|
|
||||||
});
|
|
||||||
handles.push(handle);
|
|
||||||
}
|
|
||||||
|
|
||||||
for handle in handles {
|
|
||||||
handle.join().unwrap();
|
|
||||||
}
|
|
||||||
|
|
||||||
assert_eq!(collector.get_all_metrics().len(), 10);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_metrics_collector_get_all() {
|
|
||||||
let collector = MetricsCollector::new();
|
|
||||||
collector.record_request("agent1", true, 100.0, None);
|
|
||||||
collector.record_request("agent2", true, 150.0, None);
|
|
||||||
|
|
||||||
let all = collector.get_all_metrics();
|
|
||||||
assert_eq!(all.len(), 2);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_metrics_read_while_other_writes() {
|
|
||||||
let collector = std::sync::Arc::new(MetricsCollector::new());
|
|
||||||
collector.record_request("agent1", true, 100.0, None);
|
|
||||||
|
|
||||||
let c1 = collector.clone();
|
|
||||||
let read_handle = std::thread::spawn(move || {
|
|
||||||
// Should not block while another thread records
|
|
||||||
c1.get_metrics("agent1")
|
|
||||||
});
|
|
||||||
|
|
||||||
let c2 = collector.clone();
|
|
||||||
let write_handle = std::thread::spawn(move || {
|
|
||||||
c2.record_request("agent2", true, 150.0, None);
|
|
||||||
});
|
|
||||||
|
|
||||||
read_handle.join().unwrap();
|
|
||||||
write_handle.join().unwrap();
|
|
||||||
assert_eq!(collector.get_all_metrics().len(), 2);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_metrics_collector_reset() {
|
|
||||||
let collector = MetricsCollector::new();
|
|
||||||
collector.record_request("agent1", true, 100.0, None);
|
|
||||||
assert!(collector.get_metrics("agent1").is_some());
|
|
||||||
|
|
||||||
collector.reset("agent1");
|
|
||||||
assert!(collector.get_metrics("agent1").is_none());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_metrics_isolation() {
|
|
||||||
let collector = MetricsCollector::new();
|
|
||||||
collector.record_request("agent1", true, 100.0, None);
|
|
||||||
collector.record_request("agent2", true, 150.0, None);
|
|
||||||
|
|
||||||
let m1 = collector.get_metrics("agent1").unwrap();
|
|
||||||
let m2 = collector.get_metrics("agent2").unwrap();
|
|
||||||
|
|
||||||
assert_ne!(m1.agent_id, m2.agent_id);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_latency_percentiles() {
|
|
||||||
let collector = MetricsCollector::new();
|
|
||||||
for i in 1..=30 {
|
|
||||||
collector.record_request("agent1", true, (i * 10) as f32, None);
|
|
||||||
}
|
|
||||||
|
|
||||||
let metrics = collector.get_metrics("agent1");
|
|
||||||
let m = metrics.unwrap();
|
|
||||||
assert!(m.average_latency_ms > 0.0);
|
|
||||||
assert!(m.p95_latency_ms > m.average_latency_ms);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_rwlock_behavior() {
|
|
||||||
let collector = MetricsCollector::new();
|
|
||||||
collector.record_request("agent1", true, 100.0, None);
|
|
||||||
let m1 = collector.get_metrics("agent1");
|
|
||||||
let m2 = collector.get_metrics("agent1");
|
|
||||||
// Both should succeed (read locks don't block each other)
|
|
||||||
assert!(m1.is_some());
|
|
||||||
assert!(m2.is_some());
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -224,9 +224,20 @@ impl KvCacheAligner {
|
|||||||
|
|
||||||
/// Pre-load hot chunks into cache
|
/// Pre-load hot chunks into cache
|
||||||
pub fn preload_hot_chunks(&self, hot_chunks: Vec<(&str, &str)>) -> Result<()> {
|
pub fn preload_hot_chunks(&self, hot_chunks: Vec<(&str, &str)>) -> Result<()> {
|
||||||
|
let count = hot_chunks.len();
|
||||||
for (chunk_id, text) in hot_chunks {
|
for (chunk_id, text) in hot_chunks {
|
||||||
self.cache.put(chunk_id, text);
|
self.cache.put(chunk_id, text);
|
||||||
}
|
}
|
||||||
|
let metrics = self.cache.metrics();
|
||||||
|
tracing::info!(
|
||||||
|
target: "observability",
|
||||||
|
event = "cache_preload",
|
||||||
|
preloaded = count,
|
||||||
|
cache_hits = metrics.hits,
|
||||||
|
cache_misses = metrics.misses,
|
||||||
|
hit_ratio = format!("{:.2}", metrics.hit_ratio()),
|
||||||
|
"Cache preload complete"
|
||||||
|
);
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -211,17 +211,32 @@ impl ChunkOptimizer {
|
|||||||
|
|
||||||
/// End-to-end optimization pipeline
|
/// End-to-end optimization pipeline
|
||||||
pub fn optimize(&self, chunks: Vec<OptimizableChunk>) -> (Vec<OptimizableChunk>, SelectionMetrics) {
|
pub fn optimize(&self, chunks: Vec<OptimizableChunk>) -> (Vec<OptimizableChunk>, SelectionMetrics) {
|
||||||
|
let input_count = chunks.len();
|
||||||
|
|
||||||
// Step 1: Filter by threshold
|
// Step 1: Filter by threshold
|
||||||
let filtered = self.threshold_filter.filter(chunks.clone());
|
let filtered = self.threshold_filter.filter(chunks.clone());
|
||||||
|
let after_filter = filtered.len();
|
||||||
|
|
||||||
// Step 2: Deduplicate
|
// Step 2: Deduplicate
|
||||||
let (deduplicated, dedup_removed) = self.deduplicator.deduplicate(filtered);
|
let (deduplicated, dedup_removed) = self.deduplicator.deduplicate(filtered);
|
||||||
|
let after_dedup = deduplicated.len();
|
||||||
|
|
||||||
// Step 3: Select within budget
|
// Step 3: Select within budget
|
||||||
let (selected, mut metrics) = self.budget_selector.select(deduplicated);
|
let (selected, mut metrics) = self.budget_selector.select(deduplicated);
|
||||||
|
|
||||||
metrics.dedup_removed = dedup_removed;
|
metrics.dedup_removed = dedup_removed;
|
||||||
|
|
||||||
|
tracing::info!(
|
||||||
|
target: "observability",
|
||||||
|
event = "chunk_optimize",
|
||||||
|
input = input_count,
|
||||||
|
after_threshold_filter = after_filter,
|
||||||
|
after_dedup = after_dedup,
|
||||||
|
dedup_removed = dedup_removed,
|
||||||
|
selected = selected.len(),
|
||||||
|
budget_bytes = metrics.total_bytes,
|
||||||
|
"Chunk optimization complete"
|
||||||
|
);
|
||||||
|
|
||||||
(selected, metrics)
|
(selected, metrics)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -346,7 +346,19 @@ pub async fn compact_memory(
|
|||||||
}
|
}
|
||||||
|
|
||||||
total_stats.duration_ms = start.elapsed().as_millis() as u64;
|
total_stats.duration_ms = start.elapsed().as_millis() as u64;
|
||||||
info!("Compaction complete in {}ms: {:?}", total_stats.duration_ms, total_stats);
|
info!(
|
||||||
|
target: "observability",
|
||||||
|
event = "compaction_complete",
|
||||||
|
mode = ?mode,
|
||||||
|
duration_ms = total_stats.duration_ms,
|
||||||
|
duplicate_edges_deleted = total_stats.duplicate_edges_deleted,
|
||||||
|
stale_facts_deleted = total_stats.stale_facts_deleted,
|
||||||
|
semantic_merged = total_stats.semantic_merged,
|
||||||
|
llm_calls = total_stats.llm_calls,
|
||||||
|
bytes_freed = total_stats.bytes_freed,
|
||||||
|
human_reviews_queued = total_stats.human_reviews_queued,
|
||||||
|
"Compaction complete"
|
||||||
|
);
|
||||||
|
|
||||||
Ok(total_stats)
|
Ok(total_stats)
|
||||||
}
|
}
|
||||||
@@ -373,6 +385,7 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
|
#[ignore = "not yet implemented - needs mock pool"]
|
||||||
fn test_confidence_thresholds() {
|
fn test_confidence_thresholds() {
|
||||||
let tier2 = Tier2Compactor::new(
|
let tier2 = Tier2Compactor::new(
|
||||||
// Mock pool would go here
|
// Mock pool would go here
|
||||||
|
|||||||
@@ -344,6 +344,22 @@ impl FullPipeline {
|
|||||||
|
|
||||||
metrics.total_latency_ms = start.elapsed().as_millis() as u64;
|
metrics.total_latency_ms = start.elapsed().as_millis() as u64;
|
||||||
|
|
||||||
|
tracing::info!(
|
||||||
|
target: "observability",
|
||||||
|
event = "full_pipeline_complete",
|
||||||
|
query = query,
|
||||||
|
candidates = metrics.wiki_scope_docs,
|
||||||
|
prefiltered = metrics.prefilter_candidates,
|
||||||
|
optimized = metrics.post_optimization_count,
|
||||||
|
dedup_removed = metrics.dedup_removed,
|
||||||
|
boosts_applied = metrics.metadata_boosts_applied,
|
||||||
|
cache_hit_ratio = format!("{:.2}", metrics.cache_hit_ratio),
|
||||||
|
budget_bytes = metrics.budget_used_bytes,
|
||||||
|
total_ms = metrics.total_latency_ms,
|
||||||
|
"Full query pipeline complete"
|
||||||
|
);
|
||||||
|
|
||||||
|
|
||||||
Ok(PipelineResult {
|
Ok(PipelineResult {
|
||||||
query: query.to_string(),
|
query: query.to_string(),
|
||||||
query_intent,
|
query_intent,
|
||||||
@@ -467,6 +483,22 @@ impl FullPipeline {
|
|||||||
|
|
||||||
metrics.total_latency_ms = start.elapsed().as_millis() as u64;
|
metrics.total_latency_ms = start.elapsed().as_millis() as u64;
|
||||||
|
|
||||||
|
tracing::info!(
|
||||||
|
target: "observability",
|
||||||
|
event = "full_pipeline_complete",
|
||||||
|
query = query,
|
||||||
|
candidates = metrics.wiki_scope_docs,
|
||||||
|
prefiltered = metrics.prefilter_candidates,
|
||||||
|
optimized = metrics.post_optimization_count,
|
||||||
|
dedup_removed = metrics.dedup_removed,
|
||||||
|
boosts_applied = metrics.metadata_boosts_applied,
|
||||||
|
cache_hit_ratio = format!("{:.2}", metrics.cache_hit_ratio),
|
||||||
|
budget_bytes = metrics.budget_used_bytes,
|
||||||
|
total_ms = metrics.total_latency_ms,
|
||||||
|
"Full query pipeline complete"
|
||||||
|
);
|
||||||
|
|
||||||
|
|
||||||
Ok(PipelineResult {
|
Ok(PipelineResult {
|
||||||
query: query.to_string(),
|
query: query.to_string(),
|
||||||
query_intent,
|
query_intent,
|
||||||
|
|||||||
@@ -317,109 +317,3 @@ pub async fn delete_agent_handler(
|
|||||||
}))
|
}))
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_register_agent_request() {
|
|
||||||
let req = RegisterAgentRequest {
|
|
||||||
agent_id: "agent1".to_string(),
|
|
||||||
project_id: "proj1".to_string(),
|
|
||||||
capabilities: vec!["summarization".to_string()],
|
|
||||||
webhook_url: None,
|
|
||||||
rate_limit: Some(500),
|
|
||||||
};
|
|
||||||
assert_eq!(req.agent_id, "agent1");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_agent_response() {
|
|
||||||
let resp = AgentResponse {
|
|
||||||
agent_id: "a1".to_string(),
|
|
||||||
project_id: "p1".to_string(),
|
|
||||||
capabilities: vec!["summarization".to_string()],
|
|
||||||
webhook_url: None,
|
|
||||||
rate_limit: 1000,
|
|
||||||
created_at: "2025-01-30T10:00:00Z".to_string(),
|
|
||||||
status: "active".to_string(),
|
|
||||||
};
|
|
||||||
assert_eq!(resp.status, "active");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_metrics_response() {
|
|
||||||
let metrics = MetricsResponse {
|
|
||||||
agent_id: "a1".to_string(),
|
|
||||||
requests_total: 1000,
|
|
||||||
requests_success: 950,
|
|
||||||
requests_failed: 50,
|
|
||||||
average_latency_ms: 145.5,
|
|
||||||
p95_latency_ms: 310.0,
|
|
||||||
p99_latency_ms: 450.0,
|
|
||||||
error_rate: 0.05,
|
|
||||||
};
|
|
||||||
assert!(metrics.error_rate < 0.1);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_update_agent_request() {
|
|
||||||
let req = UpdateAgentRequest {
|
|
||||||
webhook_url: Some("http://localhost".to_string()),
|
|
||||||
rate_limit: Some(500),
|
|
||||||
capabilities: None,
|
|
||||||
};
|
|
||||||
assert!(req.webhook_url.is_some());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_extract_jwt_token_valid() {
|
|
||||||
// Note: requires actix_web test setup - stub test
|
|
||||||
let jwt = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9";
|
|
||||||
let auth_header = format!("Bearer {}", jwt);
|
|
||||||
assert!(auth_header.starts_with("Bearer "));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_jwt_propagation_to_synthesis() {
|
|
||||||
let jwt = "test-jwt-token".to_string();
|
|
||||||
let client = SynthesisClient::new(
|
|
||||||
"http://api.riotpiao.com".to_string(),
|
|
||||||
jwt.clone(),
|
|
||||||
);
|
|
||||||
assert_eq!(client.jwt_token, jwt);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_agent_reasoning_with_same_jwt() {
|
|
||||||
let jwt = "shared-jwt-token".to_string();
|
|
||||||
let client = SynthesisClient::new(
|
|
||||||
"http://api.riotpiao.com".to_string(),
|
|
||||||
jwt.clone(),
|
|
||||||
);
|
|
||||||
assert_eq!(client.jwt_token, jwt);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_jwt_required_for_delete() {
|
|
||||||
// Deletion requires authentication via JWT token
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_synthesis_client_api_riotpiao() {
|
|
||||||
let jwt = "test-jwt".to_string();
|
|
||||||
let client = SynthesisClient::new(
|
|
||||||
"https://api.riotpiao.com".to_string(),
|
|
||||||
jwt.clone(),
|
|
||||||
);
|
|
||||||
assert!(client.base_url.contains("riotpiao"));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// QUALITY IMPROVEMENTS (Phase 6 JWT Auth):
|
|
||||||
// - extract_jwt_token() centralizes Bearer token extraction
|
|
||||||
// - All agent handlers extract and validate JWT
|
|
||||||
// - SynthesisClient receives JWT and uses for all reasoning calls
|
|
||||||
// - Consistent security context across ingest pipeline
|
|
||||||
// - Logging tracks JWT auth presence/absence
|
|
||||||
// - Deletion requires JWT (higher security)
|
|
||||||
|
|||||||
@@ -407,169 +407,3 @@ pub async fn hybrid_search_handler(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_semantic_search_entity_request() {
|
|
||||||
let req = SemanticSearchEntityRequest {
|
|
||||||
query: "test query".to_string(),
|
|
||||||
entity_type: Some("concept".to_string()),
|
|
||||||
confidence_floor: 0.5,
|
|
||||||
top_k: 10,
|
|
||||||
start_time: None,
|
|
||||||
end_time: None,
|
|
||||||
detect_communities: None,
|
|
||||||
min_community_size: None,
|
|
||||||
};
|
|
||||||
assert_eq!(req.query, "test query");
|
|
||||||
assert_eq!(req.confidence_floor, 0.5);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_semantic_search_with_temporal_range() {
|
|
||||||
use chrono::{Utc, Duration};
|
|
||||||
let now = Utc::now();
|
|
||||||
let tomorrow = now + Duration::days(1);
|
|
||||||
|
|
||||||
let req = SemanticSearchEntityRequest {
|
|
||||||
query: "test query".to_string(),
|
|
||||||
entity_type: None,
|
|
||||||
confidence_floor: 0.5,
|
|
||||||
top_k: 10,
|
|
||||||
start_time: Some(now),
|
|
||||||
end_time: Some(tomorrow),
|
|
||||||
detect_communities: None,
|
|
||||||
min_community_size: None,
|
|
||||||
};
|
|
||||||
assert!(req.start_time <= req.end_time);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_semantic_search_with_community_detection() {
|
|
||||||
let req = SemanticSearchEntityRequest {
|
|
||||||
query: "test query".to_string(),
|
|
||||||
entity_type: None,
|
|
||||||
confidence_floor: 0.5,
|
|
||||||
top_k: 10,
|
|
||||||
start_time: None,
|
|
||||||
end_time: None,
|
|
||||||
detect_communities: Some(true),
|
|
||||||
min_community_size: Some(3),
|
|
||||||
};
|
|
||||||
assert_eq!(req.detect_communities, Some(true));
|
|
||||||
assert_eq!(req.min_community_size, Some(3));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_semantic_search_edge_request() {
|
|
||||||
let req = SemanticSearchEdgeRequest {
|
|
||||||
query: "test query".to_string(),
|
|
||||||
relation_type: Some("related_to".to_string()),
|
|
||||||
top_k: 10,
|
|
||||||
start_time: None,
|
|
||||||
end_time: None,
|
|
||||||
};
|
|
||||||
assert_eq!(req.query, "test query");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_hybrid_search_request_defaults() {
|
|
||||||
let req = HybridSearchRequest {
|
|
||||||
query: "test".to_string(),
|
|
||||||
semantic_weight: default_semantic_weight(),
|
|
||||||
lexical_weight: default_lexical_weight(),
|
|
||||||
top_k: default_top_k(),
|
|
||||||
};
|
|
||||||
assert_eq!(req.semantic_weight, 0.6);
|
|
||||||
assert_eq!(req.lexical_weight, 0.4);
|
|
||||||
assert_eq!(req.top_k, 10);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_semantic_search_response() {
|
|
||||||
let response: SemanticSearchResponse<EntityResult> = SemanticSearchResponse {
|
|
||||||
query: "test".to_string(),
|
|
||||||
results: vec![],
|
|
||||||
total_count: 0,
|
|
||||||
search_time_ms: 100,
|
|
||||||
communities: None,
|
|
||||||
paths: None,
|
|
||||||
available_facets: None,
|
|
||||||
};
|
|
||||||
assert_eq!(response.query, "test");
|
|
||||||
assert_eq!(response.total_count, 0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_semantic_search_with_path_finding() {
|
|
||||||
let req = SemanticSearchEntityRequest {
|
|
||||||
query: "test query".to_string(),
|
|
||||||
entity_type: None,
|
|
||||||
confidence_floor: 0.5,
|
|
||||||
top_k: 10,
|
|
||||||
start_time: None,
|
|
||||||
end_time: None,
|
|
||||||
detect_communities: None,
|
|
||||||
min_community_size: None,
|
|
||||||
find_paths: Some(true),
|
|
||||||
target_entity_id: Some("e5".to_string()),
|
|
||||||
max_path_depth: Some(5),
|
|
||||||
k_hops: None,
|
|
||||||
facet_filters: None,
|
|
||||||
discover_facets: None,
|
|
||||||
};
|
|
||||||
assert_eq!(req.find_paths, Some(true));
|
|
||||||
assert_eq!(req.target_entity_id, Some("e5".to_string()));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_semantic_search_with_facet_discovery() {
|
|
||||||
let req = SemanticSearchEntityRequest {
|
|
||||||
query: "kubernetes".to_string(),
|
|
||||||
entity_type: None,
|
|
||||||
confidence_floor: 0.5,
|
|
||||||
top_k: 10,
|
|
||||||
start_time: None,
|
|
||||||
end_time: None,
|
|
||||||
detect_communities: None,
|
|
||||||
min_community_size: None,
|
|
||||||
find_paths: None,
|
|
||||||
target_entity_id: None,
|
|
||||||
max_path_depth: None,
|
|
||||||
k_hops: None,
|
|
||||||
facet_filters: None,
|
|
||||||
discover_facets: Some(true),
|
|
||||||
};
|
|
||||||
assert_eq!(req.discover_facets, Some(true));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_semantic_search_with_facet_filters() {
|
|
||||||
let filters = FacetFilters {
|
|
||||||
entity_types: Some(vec!["concept".to_string()]),
|
|
||||||
relation_types: None,
|
|
||||||
confidence_level: Some("high".to_string()),
|
|
||||||
date_range: None,
|
|
||||||
};
|
|
||||||
let req = SemanticSearchEntityRequest {
|
|
||||||
query: "test".to_string(),
|
|
||||||
entity_type: None,
|
|
||||||
confidence_floor: 0.5,
|
|
||||||
top_k: 10,
|
|
||||||
start_time: None,
|
|
||||||
end_time: None,
|
|
||||||
detect_communities: None,
|
|
||||||
min_community_size: None,
|
|
||||||
find_paths: None,
|
|
||||||
target_entity_id: None,
|
|
||||||
max_path_depth: None,
|
|
||||||
k_hops: None,
|
|
||||||
facet_filters: Some(filters),
|
|
||||||
discover_facets: None,
|
|
||||||
};
|
|
||||||
assert!(req.facet_filters.is_some());
|
|
||||||
assert_eq!(req.facet_filters.unwrap().confidence_level, Some("high".to_string()));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -732,129 +732,3 @@ pub async fn summarize_handler(
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_link_entities_request() {
|
|
||||||
let req = LinkEntitiesRequest {
|
|
||||||
project: "poimen".to_string(),
|
|
||||||
text: "Kubernetes is a container orchestrator.".to_string(),
|
|
||||||
};
|
|
||||||
assert_eq!(req.project, "poimen");
|
|
||||||
assert!(!req.text.is_empty());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_detect_aliases_request() {
|
|
||||||
let req = DetectAliasesRequest {
|
|
||||||
project: "poimen".to_string(),
|
|
||||||
entity_id: "e1".to_string(),
|
|
||||||
entity_name: "Kubernetes".to_string(),
|
|
||||||
text_samples: vec!["k8s is great".to_string()],
|
|
||||||
};
|
|
||||||
assert_eq!(req.entity_name, "Kubernetes");
|
|
||||||
assert_eq!(req.text_samples.len(), 1);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_suggest_merges_request() {
|
|
||||||
let req = SuggestMergesRequest {
|
|
||||||
project: "poimen".to_string(),
|
|
||||||
similarity_threshold: 0.85,
|
|
||||||
};
|
|
||||||
assert_eq!(req.similarity_threshold, 0.85);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_suggest_merges_default_threshold() {
|
|
||||||
let req = SuggestMergesRequest {
|
|
||||||
project: "poimen".to_string(),
|
|
||||||
similarity_threshold: default_merge_threshold(),
|
|
||||||
};
|
|
||||||
assert_eq!(req.similarity_threshold, 0.8);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_detect_coreferences_request() {
|
|
||||||
let req = DetectCoreferencesRequest {
|
|
||||||
project: "poimen".to_string(),
|
|
||||||
texts: vec![
|
|
||||||
"Kubernetes is great.".to_string(),
|
|
||||||
"k8s makes deployments easy.".to_string(),
|
|
||||||
],
|
|
||||||
};
|
|
||||||
assert_eq!(req.texts.len(), 2);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_link_entities_response() {
|
|
||||||
let resp = LinkEntitiesResponse {
|
|
||||||
links: vec![],
|
|
||||||
unlinked: vec![],
|
|
||||||
total_mentions: 0,
|
|
||||||
link_rate: 0.0,
|
|
||||||
process_time_ms: 100,
|
|
||||||
};
|
|
||||||
assert_eq!(resp.total_mentions, 0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_detect_aliases_response() {
|
|
||||||
let resp = DetectAliasesResponse {
|
|
||||||
entity_id: "e1".to_string(),
|
|
||||||
entity_name: "Kubernetes".to_string(),
|
|
||||||
aliases: vec![],
|
|
||||||
alias_count: 0,
|
|
||||||
process_time_ms: 100,
|
|
||||||
};
|
|
||||||
assert_eq!(resp.alias_count, 0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_suggest_merges_response() {
|
|
||||||
let resp = SuggestMergesResponse {
|
|
||||||
project: "poimen".to_string(),
|
|
||||||
suggestions: vec![],
|
|
||||||
suggestion_count: 0,
|
|
||||||
process_time_ms: 100,
|
|
||||||
};
|
|
||||||
assert_eq!(resp.suggestion_count, 0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_detect_coreferences_response() {
|
|
||||||
let resp = DetectCoreferencesResponse {
|
|
||||||
project: "poimen".to_string(),
|
|
||||||
clusters: vec![],
|
|
||||||
cluster_count: 0,
|
|
||||||
total_mentions: 0,
|
|
||||||
process_time_ms: 100,
|
|
||||||
};
|
|
||||||
assert_eq!(resp.cluster_count, 0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_link_entities_request_serialization() {
|
|
||||||
let req = LinkEntitiesRequest {
|
|
||||||
project: "test".to_string(),
|
|
||||||
text: "Kubernetes".to_string(),
|
|
||||||
};
|
|
||||||
let json = serde_json::to_string(&req).unwrap();
|
|
||||||
assert!(json.contains("test"));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_link_entities_response_serialization() {
|
|
||||||
let resp = LinkEntitiesResponse {
|
|
||||||
links: vec![],
|
|
||||||
unlinked: vec![],
|
|
||||||
total_mentions: 5,
|
|
||||||
link_rate: 0.8,
|
|
||||||
process_time_ms: 150,
|
|
||||||
};
|
|
||||||
let json = serde_json::to_string(&resp).unwrap();
|
|
||||||
assert!(json.contains("0.8"));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -2,14 +2,15 @@ use anyhow::Result;
|
|||||||
use mem_store::{MemoryL1, VectorStore, ChunkL0, EntityRepoOps, EdgeRepoOps};
|
use mem_store::{MemoryL1, VectorStore, ChunkL0, EntityRepoOps, EdgeRepoOps};
|
||||||
use mem_llm::EmbeddingsClient;
|
use mem_llm::EmbeddingsClient;
|
||||||
use mem_ingest::ingest_pipeline::{IngestPipeline, Episode};
|
use mem_ingest::ingest_pipeline::{IngestPipeline, Episode};
|
||||||
use mem_ingest::entity_extractor::WikiLinkFallbackExtractor;
|
use mem_ingest::entity_extractor::{WikiLinkFallbackExtractor, LlmEntityExtractor};
|
||||||
use mem_ingest::fact_extractor::SimpleFactExtractor;
|
use mem_ingest::fact_extractor::{SimpleFactExtractor, LlmFactExtractor};
|
||||||
use mem_ingest::contradiction_detector::ContradictionHandler;
|
use mem_ingest::contradiction_detector::ContradictionHandler;
|
||||||
use sqlx::PgPool;
|
use sqlx::PgPool;
|
||||||
use uuid::Uuid;
|
use uuid::Uuid;
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
use pgvector::Vector;
|
use pgvector::Vector;
|
||||||
|
|
||||||
|
|
||||||
/// Ingest worker — processes queued records through entity/fact extraction pipeline
|
/// Ingest worker — processes queued records through entity/fact extraction pipeline
|
||||||
pub struct IngestWorker {
|
pub struct IngestWorker {
|
||||||
pool: PgPool,
|
pool: PgPool,
|
||||||
@@ -26,11 +27,25 @@ impl IngestWorker {
|
|||||||
) -> Self {
|
) -> Self {
|
||||||
let vector_store = Arc::new(VectorStore::new(pool.clone()));
|
let vector_store = Arc::new(VectorStore::new(pool.clone()));
|
||||||
|
|
||||||
// Initialize extraction pipeline
|
// Initialize extraction pipeline — use LLM if LLM_ENDPOINT is set, else fallback to wiki links
|
||||||
let entity_extractor: Arc<dyn mem_ingest::entity_extractor::EntityExtractor> =
|
let entity_extractor: Arc<dyn mem_ingest::entity_extractor::EntityExtractor> =
|
||||||
Arc::new(WikiLinkFallbackExtractor);
|
if std::env::var("LLM_ENDPOINT").is_ok() {
|
||||||
|
let model = std::env::var("LLM_MODEL").unwrap_or_else(|_| "qwen2.5:3b-instruct".to_string());
|
||||||
|
tracing::info!("Using LLM entity extractor: model={}", model);
|
||||||
|
Arc::new(LlmEntityExtractor::new(&model))
|
||||||
|
} else {
|
||||||
|
tracing::info!("LLM_ENDPOINT not set, using WikiLink fallback extractor");
|
||||||
|
Arc::new(WikiLinkFallbackExtractor)
|
||||||
|
};
|
||||||
let fact_extractor: Arc<dyn mem_ingest::fact_extractor::FactExtractor> =
|
let fact_extractor: Arc<dyn mem_ingest::fact_extractor::FactExtractor> =
|
||||||
Arc::new(SimpleFactExtractor);
|
if std::env::var("LLM_ENDPOINT").is_ok() {
|
||||||
|
let model = std::env::var("LLM_MODEL").unwrap_or_else(|_| "qwen2.5:3b-instruct".to_string());
|
||||||
|
tracing::info!("Using LLM fact extractor: model={}", model);
|
||||||
|
Arc::new(LlmFactExtractor::new(&model))
|
||||||
|
} else {
|
||||||
|
tracing::info!("LLM_ENDPOINT not set, using simple pattern fact extractor");
|
||||||
|
Arc::new(SimpleFactExtractor)
|
||||||
|
};
|
||||||
let contradiction_detector = Arc::new(ContradictionHandler::default());
|
let contradiction_detector = Arc::new(ContradictionHandler::default());
|
||||||
let pipeline = Arc::new(IngestPipeline::new(
|
let pipeline = Arc::new(IngestPipeline::new(
|
||||||
entity_extractor,
|
entity_extractor,
|
||||||
@@ -121,9 +136,15 @@ impl IngestWorker {
|
|||||||
.await?;
|
.await?;
|
||||||
|
|
||||||
tracing::info!(
|
tracing::info!(
|
||||||
"Ingest completed: {} (entities={}, edges={}, reviews={})",
|
target: "observability",
|
||||||
ingest_id, total_entities, total_edges, total_reviews
|
event = "ingest_complete",
|
||||||
|
ingest_id = ingest_id,
|
||||||
|
entities = total_entities,
|
||||||
|
edges = total_edges,
|
||||||
|
reviews = total_reviews,
|
||||||
|
"Ingest completed"
|
||||||
);
|
);
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -173,7 +194,12 @@ async fn save_entity_to_db(pool: &PgPool, entity: &mem_core::entity::Entity) ->
|
|||||||
sqlx::query(
|
sqlx::query(
|
||||||
"INSERT INTO memory_entity (id, project_id, name, entity_type, description, t_created, t_updated, confidence)
|
"INSERT INTO memory_entity (id, project_id, name, entity_type, description, t_created, t_updated, confidence)
|
||||||
VALUES ($1, $2, $3, $4, $5, $6::TIMESTAMPTZ, $7::TIMESTAMPTZ, $8)
|
VALUES ($1, $2, $3, $4, $5, $6::TIMESTAMPTZ, $7::TIMESTAMPTZ, $8)
|
||||||
ON CONFLICT (id) DO NOTHING"
|
ON CONFLICT (project_id, name) DO UPDATE SET
|
||||||
|
entity_type = EXCLUDED.entity_type,
|
||||||
|
description = COALESCE(NULLIF(EXCLUDED.description, ''), memory_entity.description),
|
||||||
|
t_updated = NOW(),
|
||||||
|
confidence = GREATEST(memory_entity.confidence, EXCLUDED.confidence),
|
||||||
|
source_count = memory_entity.source_count + 1"
|
||||||
)
|
)
|
||||||
.bind(&entity.id)
|
.bind(&entity.id)
|
||||||
.bind(&entity.project_id)
|
.bind(&entity.project_id)
|
||||||
@@ -193,7 +219,7 @@ async fn save_entity_to_db(pool: &PgPool, entity: &mem_core::entity::Entity) ->
|
|||||||
async fn save_edge_to_db(pool: &PgPool, edge: &mem_core::edge::Edge) -> Result<()> {
|
async fn save_edge_to_db(pool: &PgPool, edge: &mem_core::edge::Edge) -> Result<()> {
|
||||||
// Try temporal schema first (id, project_id, source_entity_id, etc)
|
// Try temporal schema first (id, project_id, source_entity_id, etc)
|
||||||
let result = sqlx::query(
|
let result = sqlx::query(
|
||||||
"INSERT INTO memory_edge (id, project_id, source_entity_id, target_entity_id, relation_type, fact, t_valid, t_invalid, t_created, confidence)
|
"INSERT INTO memory_edge (id, project_id, source_id, target_id, relation_type, fact, t_valid, t_invalid, t_created, confidence)
|
||||||
VALUES ($1, $2, $3, $4, $5, $6, $7::TIMESTAMPTZ, $8::TIMESTAMPTZ, $9::TIMESTAMPTZ, $10)
|
VALUES ($1, $2, $3, $4, $5, $6, $7::TIMESTAMPTZ, $8::TIMESTAMPTZ, $9::TIMESTAMPTZ, $10)
|
||||||
ON CONFLICT (id) DO NOTHING"
|
ON CONFLICT (id) DO NOTHING"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -158,106 +158,3 @@ impl ParallelDualWriteIndexer {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_indexable_chunk_structure() {
|
|
||||||
let chunk = IndexableChunk {
|
|
||||||
chunk_id: "c1".to_string(),
|
|
||||||
content: "test".to_string(),
|
|
||||||
source: "src".to_string(),
|
|
||||||
project: "proj".to_string(),
|
|
||||||
level: "L1".to_string(),
|
|
||||||
breadcrumb: vec!["a".to_string()],
|
|
||||||
};
|
|
||||||
assert_eq!(chunk.chunk_id, "c1");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_dual_write_result_structure() {
|
|
||||||
let result = DualWriteResult {
|
|
||||||
chunk_id: "c1".to_string(),
|
|
||||||
pgvector_success: true,
|
|
||||||
opensearch_success: true,
|
|
||||||
error: None,
|
|
||||||
};
|
|
||||||
assert!(result.pgvector_success);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_parallel_indexer_creation() {
|
|
||||||
let pool = sqlx::postgres::PgPoolOptions::new()
|
|
||||||
.max_connections(1)
|
|
||||||
.build_lazy();
|
|
||||||
let indexer = ParallelDualWriteIndexer::new(pool, None);
|
|
||||||
assert!(indexer.opensearch.is_none());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_hash_computation() {
|
|
||||||
let pool = sqlx::postgres::PgPoolOptions::new()
|
|
||||||
.max_connections(1)
|
|
||||||
.build_lazy();
|
|
||||||
let indexer = ParallelDualWriteIndexer::new(pool, None);
|
|
||||||
let hash1 = indexer.compute_hash("test");
|
|
||||||
let hash2 = indexer.compute_hash("test");
|
|
||||||
assert_eq!(hash1, hash2);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_hash_different_content() {
|
|
||||||
let pool = sqlx::postgres::PgPoolOptions::new()
|
|
||||||
.max_connections(1)
|
|
||||||
.build_lazy();
|
|
||||||
let indexer = ParallelDualWriteIndexer::new(pool, None);
|
|
||||||
let hash1 = indexer.compute_hash("test1");
|
|
||||||
let hash2 = indexer.compute_hash("test2");
|
|
||||||
assert_ne!(hash1, hash2);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_dual_write_result_pgvector_failed() {
|
|
||||||
let result = DualWriteResult {
|
|
||||||
chunk_id: "c1".to_string(),
|
|
||||||
pgvector_success: false,
|
|
||||||
opensearch_success: true,
|
|
||||||
error: Some("pgvector failed".to_string()),
|
|
||||||
};
|
|
||||||
assert!(!result.pgvector_success);
|
|
||||||
assert!(result.error.is_some());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_dual_write_result_opensearch_failed() {
|
|
||||||
let result = DualWriteResult {
|
|
||||||
chunk_id: "c1".to_string(),
|
|
||||||
pgvector_success: true,
|
|
||||||
opensearch_success: false,
|
|
||||||
error: Some("opensearch failed".to_string()),
|
|
||||||
};
|
|
||||||
assert!(result.pgvector_success);
|
|
||||||
assert!(!result.opensearch_success);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_breadcrumb_join() {
|
|
||||||
let breadcrumb = vec!["a".to_string(), "b".to_string(), "c".to_string()];
|
|
||||||
let joined = breadcrumb.join(" > ");
|
|
||||||
assert_eq!(joined, "a > b > c");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_chunk_source_tracking() {
|
|
||||||
let chunk = IndexableChunk {
|
|
||||||
chunk_id: "c1".to_string(),
|
|
||||||
content: "test".to_string(),
|
|
||||||
source: "transcript://session-123".to_string(),
|
|
||||||
project: "poimen".to_string(),
|
|
||||||
level: "L1".to_string(),
|
|
||||||
breadcrumb: vec![],
|
|
||||||
};
|
|
||||||
assert!(chunk.source.contains("session"));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -285,11 +285,9 @@ impl BfsGraphTraversal {
|
|||||||
pub fn truncate_to_depth(graph: &mut GraphData, max_depth: i32) {
|
pub fn truncate_to_depth(graph: &mut GraphData, max_depth: i32) {
|
||||||
graph.nodes.retain(|n| n.depth <= max_depth);
|
graph.nodes.retain(|n| n.depth <= max_depth);
|
||||||
graph.edges.retain(|e| {
|
graph.edges.retain(|e| {
|
||||||
let source_depth = graph.nodes.iter()
|
let source_exists = graph.nodes.iter().any(|n| n.id == e.source_id);
|
||||||
.find(|n| n.id == e.source_id)
|
let target_exists = graph.nodes.iter().any(|n| n.id == e.target_id);
|
||||||
.map(|n| n.depth)
|
source_exists && target_exists
|
||||||
.unwrap_or(i32::MAX);
|
|
||||||
source_depth <= max_depth
|
|
||||||
});
|
});
|
||||||
|
|
||||||
graph.max_depth_reached = graph.max_depth_reached.min(max_depth);
|
graph.max_depth_reached = graph.max_depth_reached.min(max_depth);
|
||||||
|
|||||||
@@ -343,167 +343,3 @@ impl CommunityDetector {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_community_creation() {
|
|
||||||
let community = Community {
|
|
||||||
id: 0,
|
|
||||||
entity_ids: vec!["e1".to_string(), "e2".to_string()],
|
|
||||||
entity_names: vec!["Entity1".to_string(), "Entity2".to_string()],
|
|
||||||
size: 2,
|
|
||||||
modularity_contribution: 0.8,
|
|
||||||
average_strength: 0.9,
|
|
||||||
density: 1.0,
|
|
||||||
};
|
|
||||||
assert_eq!(community.size, 2);
|
|
||||||
assert_eq!(community.entity_ids.len(), 2);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_community_detection_result() {
|
|
||||||
let result = CommunityDetectionResult {
|
|
||||||
entity_count: 100,
|
|
||||||
edge_count: 250,
|
|
||||||
communities: vec![],
|
|
||||||
community_count: 0,
|
|
||||||
total_modularity: 0.0,
|
|
||||||
average_community_size: 0.0,
|
|
||||||
};
|
|
||||||
assert_eq!(result.entity_count, 100);
|
|
||||||
assert_eq!(result.edge_count, 250);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_min_community_size_clamping() {
|
|
||||||
let size = 1;
|
|
||||||
let clamped = size.max(2).min(1000);
|
|
||||||
assert_eq!(clamped, 2);
|
|
||||||
|
|
||||||
let size = 5000;
|
|
||||||
let clamped = size.max(2).min(1000);
|
|
||||||
assert_eq!(clamped, 1000);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_modularity_threshold_clamping() {
|
|
||||||
let threshold = 0.0001;
|
|
||||||
let clamped = threshold.max(0.0001).min(0.1);
|
|
||||||
assert_eq!(clamped, 0.0001);
|
|
||||||
|
|
||||||
let threshold = 0.5;
|
|
||||||
let clamped = threshold.max(0.0001).min(0.1);
|
|
||||||
assert_eq!(clamped, 0.1);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_density_calculation() {
|
|
||||||
// 3 entities, all connected (3 edges)
|
|
||||||
// Possible edges: 3 * 2 / 2 = 3
|
|
||||||
// Density: 3 / 3 = 1.0 (fully connected)
|
|
||||||
let density = (3.0 / 3.0).max(0.0).min(1.0);
|
|
||||||
assert_eq!(density, 1.0);
|
|
||||||
|
|
||||||
// 4 entities, 2 edges
|
|
||||||
// Possible: 4 * 3 / 2 = 6
|
|
||||||
// Density: 2 / 6 ≈ 0.33
|
|
||||||
let density = (2.0 / 6.0).max(0.0).min(1.0);
|
|
||||||
assert!((density - 0.333).abs() < 0.01);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_modularity_bounds() {
|
|
||||||
let modularity = 0.75;
|
|
||||||
let clamped = modularity.max(-1.0).min(1.0);
|
|
||||||
assert_eq!(clamped, 0.75);
|
|
||||||
|
|
||||||
let modularity = -0.5;
|
|
||||||
let clamped = modularity.max(-1.0).min(1.0);
|
|
||||||
assert_eq!(clamped, -0.5);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_average_community_size() {
|
|
||||||
let communities = vec![
|
|
||||||
Community {
|
|
||||||
id: 0,
|
|
||||||
entity_ids: vec!["a".into(), "b".into(), "c".into()],
|
|
||||||
entity_names: vec![],
|
|
||||||
size: 3,
|
|
||||||
modularity_contribution: 0.5,
|
|
||||||
average_strength: 0.8,
|
|
||||||
density: 0.9,
|
|
||||||
},
|
|
||||||
Community {
|
|
||||||
id: 1,
|
|
||||||
entity_ids: vec!["d".into(), "e".into()],
|
|
||||||
entity_names: vec![],
|
|
||||||
size: 2,
|
|
||||||
modularity_contribution: 0.4,
|
|
||||||
average_strength: 0.7,
|
|
||||||
density: 1.0,
|
|
||||||
},
|
|
||||||
];
|
|
||||||
|
|
||||||
let avg = communities.iter().map(|c| c.size as f32).sum::<f32>() / communities.len() as f32;
|
|
||||||
assert_eq!(avg, 2.5);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_total_modularity_sum() {
|
|
||||||
let contributions = vec![0.3, 0.25, 0.2, 0.15];
|
|
||||||
let total: f32 = contributions.iter().sum();
|
|
||||||
let clamped = total.max(-1.0).min(1.0);
|
|
||||||
|
|
||||||
assert!(clamped >= -1.0 && clamped <= 1.0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_empty_graph_handling() {
|
|
||||||
let entities: Vec<String> = vec![];
|
|
||||||
let edges: Vec<GraphEdge> = vec![];
|
|
||||||
|
|
||||||
assert!(entities.is_empty());
|
|
||||||
assert!(edges.is_empty());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_single_node_graph() {
|
|
||||||
let entity_count = 1;
|
|
||||||
let edge_count = 0;
|
|
||||||
|
|
||||||
assert_eq!(entity_count, 1);
|
|
||||||
assert_eq!(edge_count, 0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_fully_connected_graph() {
|
|
||||||
// 5 nodes fully connected: 5*4/2 = 10 edges
|
|
||||||
let nodes = 5;
|
|
||||||
let possible_edges = nodes * (nodes - 1) / 2;
|
|
||||||
assert_eq!(possible_edges, 10);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_strength_normalization() {
|
|
||||||
let strengths = vec![0.0, 0.25, 0.5, 0.75, 1.0];
|
|
||||||
for s in strengths {
|
|
||||||
let normalized = s.max(0.0).min(1.0);
|
|
||||||
assert!(normalized >= 0.0 && normalized <= 1.0);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_louvain_max_iterations() {
|
|
||||||
let max_iterations = 100;
|
|
||||||
let mut iteration = 0;
|
|
||||||
|
|
||||||
while iteration < max_iterations && iteration < 5 {
|
|
||||||
iteration += 1;
|
|
||||||
}
|
|
||||||
|
|
||||||
assert!(iteration <= max_iterations);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -437,180 +437,3 @@ struct EntityInfo {
|
|||||||
name: String,
|
name: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
fn create_linker_mock() -> EntityLinker {
|
|
||||||
// Create with in-memory pool (stub for testing)
|
|
||||||
let pool = sqlx::postgres::PgPoolOptions::new()
|
|
||||||
.max_connections(1)
|
|
||||||
.build_lazy();
|
|
||||||
EntityLinker::new(pool)
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_extract_mentions_basic() {
|
|
||||||
let linker = create_linker_mock();
|
|
||||||
let text = "Kubernetes is a container orchestration platform.";
|
|
||||||
let mentions = linker.extract_mentions(text).unwrap();
|
|
||||||
assert!(mentions.len() > 0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_extract_mentions_multiword() {
|
|
||||||
let linker = create_linker_mock();
|
|
||||||
let text = "Google Cloud Platform provides services.";
|
|
||||||
let mentions = linker.extract_mentions(text).unwrap();
|
|
||||||
assert!(mentions.iter().any(|m| m.text.contains("Cloud")));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_mention_link_structure() {
|
|
||||||
let link = MentionLink {
|
|
||||||
mention_text: "Kubernetes".to_string(),
|
|
||||||
start_offset: 0,
|
|
||||||
end_offset: 10,
|
|
||||||
entity_id: "e1".to_string(),
|
|
||||||
entity_name: "Kubernetes".to_string(),
|
|
||||||
confidence: 0.95,
|
|
||||||
reason: LinkReason::LexicalMatch,
|
|
||||||
};
|
|
||||||
assert_eq!(link.confidence, 0.95);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_link_reason_enum() {
|
|
||||||
let reasons = vec![
|
|
||||||
LinkReason::SemanticMatch,
|
|
||||||
LinkReason::LexicalMatch,
|
|
||||||
LinkReason::AliasMatch,
|
|
||||||
LinkReason::AcronymMatch,
|
|
||||||
LinkReason::PartialMatch,
|
|
||||||
];
|
|
||||||
assert_eq!(reasons.len(), 5);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_alias_suggestion_structure() {
|
|
||||||
let alias = AliasSuggestion {
|
|
||||||
entity_id: "e1".to_string(),
|
|
||||||
canonical_name: "Kubernetes".to_string(),
|
|
||||||
alias: "k8s".to_string(),
|
|
||||||
confidence: 0.9,
|
|
||||||
frequency: 5,
|
|
||||||
};
|
|
||||||
assert_eq!(alias.frequency, 5);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_merge_suggestion_structure() {
|
|
||||||
let merge = MergeSuggestion {
|
|
||||||
entity1_id: "e1".to_string(),
|
|
||||||
entity1_name: "Kubernetes".to_string(),
|
|
||||||
entity2_id: "e2".to_string(),
|
|
||||||
entity2_name: "K8s".to_string(),
|
|
||||||
confidence: 0.85,
|
|
||||||
reasons: vec!["Acronym match".to_string()],
|
|
||||||
};
|
|
||||||
assert_eq!(merge.confidence, 0.85);
|
|
||||||
assert_eq!(merge.reasons.len(), 1);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_coreference_cluster_structure() {
|
|
||||||
let cluster = CoreferenceCluster {
|
|
||||||
entity_id: "e1".to_string(),
|
|
||||||
mentions: vec!["Kubernetes".to_string(), "k8s".to_string()],
|
|
||||||
mention_count: 2,
|
|
||||||
confidence: 0.85,
|
|
||||||
};
|
|
||||||
assert_eq!(cluster.mention_count, 2);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_edit_distance() {
|
|
||||||
let linker = create_linker_mock();
|
|
||||||
let dist = linker.edit_distance("Kubernetes", "kubernetes");
|
|
||||||
assert_eq!(dist, 0); // Same lowercase
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_edit_distance_typo() {
|
|
||||||
let linker = create_linker_mock();
|
|
||||||
let dist = linker.edit_distance("Kubernetes", "Kubenetes");
|
|
||||||
assert!(dist > 0 && dist < 5);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_compute_similarity_exact() {
|
|
||||||
let linker = create_linker_mock();
|
|
||||||
let sim = linker.compute_similarity("test", "test");
|
|
||||||
assert_eq!(sim, 1.0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_compute_similarity_case_insensitive() {
|
|
||||||
let linker = create_linker_mock();
|
|
||||||
let sim = linker.compute_similarity("Test", "test");
|
|
||||||
assert_eq!(sim, 1.0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_compute_similarity_substring() {
|
|
||||||
let linker = create_linker_mock();
|
|
||||||
let sim = linker.compute_similarity("Kubernetes", "kubernetes");
|
|
||||||
assert!(sim > 0.8);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_is_acronym_true() {
|
|
||||||
let linker = create_linker_mock();
|
|
||||||
let is_acr = linker.is_acronym("k8s", "Kubernetes");
|
|
||||||
assert!(is_acr);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_is_acronym_false() {
|
|
||||||
let linker = create_linker_mock();
|
|
||||||
let is_acr = linker.is_acronym("test", "Kubernetes");
|
|
||||||
assert!(!is_acr);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_is_similar_true() {
|
|
||||||
let linker = create_linker_mock();
|
|
||||||
let similar = linker.is_similar("Kubernetes", "kubernetes");
|
|
||||||
assert!(similar);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_is_similar_false() {
|
|
||||||
let linker = create_linker_mock();
|
|
||||||
let similar = linker.is_similar("test", "completely different");
|
|
||||||
assert!(!similar);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_mention_link_reason_serialization() {
|
|
||||||
let reason = LinkReason::SemanticMatch;
|
|
||||||
let json = serde_json::to_string(&reason).unwrap();
|
|
||||||
assert!(json.contains("SemanticMatch"));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_mention_link_full_serialization() {
|
|
||||||
let link = MentionLink {
|
|
||||||
mention_text: "Kubernetes".to_string(),
|
|
||||||
start_offset: 0,
|
|
||||||
end_offset: 10,
|
|
||||||
entity_id: "e1".to_string(),
|
|
||||||
entity_name: "Kubernetes".to_string(),
|
|
||||||
confidence: 0.95,
|
|
||||||
reason: LinkReason::LexicalMatch,
|
|
||||||
};
|
|
||||||
let json = serde_json::to_string(&link).unwrap();
|
|
||||||
assert!(json.contains("Kubernetes"));
|
|
||||||
assert!(json.contains("0.95"));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -360,252 +360,3 @@ impl FacetedSearch {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_facet_value_creation() {
|
|
||||||
let facet = FacetValue {
|
|
||||||
name: "concept".to_string(),
|
|
||||||
count: 42,
|
|
||||||
percentage: 15.5,
|
|
||||||
};
|
|
||||||
|
|
||||||
assert_eq!(facet.name, "concept");
|
|
||||||
assert_eq!(facet.count, 42);
|
|
||||||
assert!((facet.percentage - 15.5).abs() < 0.01);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_facet_type_enum() {
|
|
||||||
let types = vec![
|
|
||||||
FacetType::EntityType,
|
|
||||||
FacetType::RelationType,
|
|
||||||
FacetType::ConfidenceLevel,
|
|
||||||
FacetType::DateRange,
|
|
||||||
];
|
|
||||||
|
|
||||||
assert_eq!(types.len(), 4);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_facet_filters_default() {
|
|
||||||
let filters = FacetFilters::default();
|
|
||||||
|
|
||||||
assert!(filters.entity_types.is_none());
|
|
||||||
assert!(filters.relation_types.is_none());
|
|
||||||
assert!(filters.confidence_level.is_none());
|
|
||||||
assert!(filters.date_range.is_none());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_confidence_floor_high() {
|
|
||||||
let engine = FacetedSearch { pool: unsafe { std::mem::zeroed() } };
|
|
||||||
let floor = engine.confidence_floor_from_level(Some("high"));
|
|
||||||
|
|
||||||
assert_eq!(floor, 0.8);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_confidence_floor_medium() {
|
|
||||||
let engine = FacetedSearch { pool: unsafe { std::mem::zeroed() } };
|
|
||||||
let floor = engine.confidence_floor_from_level(Some("medium"));
|
|
||||||
|
|
||||||
assert_eq!(floor, 0.5);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_confidence_floor_low() {
|
|
||||||
let engine = FacetedSearch { pool: unsafe { std::mem::zeroed() } };
|
|
||||||
let floor = engine.confidence_floor_from_level(Some("low"));
|
|
||||||
|
|
||||||
assert_eq!(floor, 0.0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_confidence_floor_none() {
|
|
||||||
let engine = FacetedSearch { pool: unsafe { std::mem::zeroed() } };
|
|
||||||
let floor = engine.confidence_floor_from_level(None);
|
|
||||||
|
|
||||||
assert_eq!(floor, 0.0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_facet_percentage_calculation() {
|
|
||||||
let count = 25;
|
|
||||||
let total = 100;
|
|
||||||
let percentage = (count as f32 / total as f32) * 100.0;
|
|
||||||
|
|
||||||
assert_eq!(percentage, 25.0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_facet_percentage_zero_total() {
|
|
||||||
let total = 0;
|
|
||||||
let percentage = if total > 0 { 100.0 } else { 0.0 };
|
|
||||||
|
|
||||||
assert_eq!(percentage, 0.0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_date_range_today() {
|
|
||||||
let engine = FacetedSearch { pool: unsafe { std::mem::zeroed() } };
|
|
||||||
let (start, end) = engine.date_range_to_times(Some("today"));
|
|
||||||
|
|
||||||
assert!(start.is_some());
|
|
||||||
assert!(end.is_some());
|
|
||||||
assert!(start.unwrap() < end.unwrap());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_date_range_week() {
|
|
||||||
let engine = FacetedSearch { pool: unsafe { std::mem::zeroed() } };
|
|
||||||
let (start, end) = engine.date_range_to_times(Some("this_week"));
|
|
||||||
|
|
||||||
assert!(start.is_some());
|
|
||||||
assert!(end.is_some());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_date_range_month() {
|
|
||||||
let engine = FacetedSearch { pool: unsafe { std::mem::zeroed() } };
|
|
||||||
let (start, end) = engine.date_range_to_times(Some("this_month"));
|
|
||||||
|
|
||||||
assert!(start.is_some());
|
|
||||||
assert!(end.is_some());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_date_range_none() {
|
|
||||||
let engine = FacetedSearch { pool: unsafe { std::mem::zeroed() } };
|
|
||||||
let (start, end) = engine.date_range_to_times(None);
|
|
||||||
|
|
||||||
assert!(start.is_none());
|
|
||||||
assert!(end.is_none());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_validate_filters_empty_entity_types() {
|
|
||||||
let engine = FacetedSearch { pool: unsafe { std::mem::zeroed() } };
|
|
||||||
let filters = FacetFilters {
|
|
||||||
entity_types: Some(vec![]),
|
|
||||||
..Default::default()
|
|
||||||
};
|
|
||||||
|
|
||||||
assert!(engine.validate_filters(&filters).is_err());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_validate_filters_valid_entity_types() {
|
|
||||||
let engine = FacetedSearch { pool: unsafe { std::mem::zeroed() } };
|
|
||||||
let filters = FacetFilters {
|
|
||||||
entity_types: Some(vec!["concept".to_string()]),
|
|
||||||
..Default::default()
|
|
||||||
};
|
|
||||||
|
|
||||||
assert!(engine.validate_filters(&filters).is_ok());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_validate_filters_too_many_types() {
|
|
||||||
let engine = FacetedSearch { pool: unsafe { std::mem::zeroed() } };
|
|
||||||
let filters = FacetFilters {
|
|
||||||
entity_types: Some((0..60).map(|i| format!("type_{}", i)).collect()),
|
|
||||||
..Default::default()
|
|
||||||
};
|
|
||||||
|
|
||||||
assert!(engine.validate_filters(&filters).is_err());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_validate_filters_invalid_confidence() {
|
|
||||||
let engine = FacetedSearch { pool: unsafe { std::mem::zeroed() } };
|
|
||||||
let filters = FacetFilters {
|
|
||||||
confidence_level: Some("invalid".to_string()),
|
|
||||||
..Default::default()
|
|
||||||
};
|
|
||||||
|
|
||||||
assert!(engine.validate_filters(&filters).is_err());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_validate_filters_valid_confidence() {
|
|
||||||
let engine = FacetedSearch { pool: unsafe { std::mem::zeroed() } };
|
|
||||||
let filters = FacetFilters {
|
|
||||||
confidence_level: Some("high".to_string()),
|
|
||||||
..Default::default()
|
|
||||||
};
|
|
||||||
|
|
||||||
assert!(engine.validate_filters(&filters).is_ok());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_validate_filters_invalid_date_range() {
|
|
||||||
let engine = FacetedSearch { pool: unsafe { std::mem::zeroed() } };
|
|
||||||
let filters = FacetFilters {
|
|
||||||
date_range: Some("invalid".to_string()),
|
|
||||||
..Default::default()
|
|
||||||
};
|
|
||||||
|
|
||||||
assert!(engine.validate_filters(&filters).is_err());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_validate_filters_valid_date_range() {
|
|
||||||
let engine = FacetedSearch { pool: unsafe { std::mem::zeroed() } };
|
|
||||||
let filters = FacetFilters {
|
|
||||||
date_range: Some("this_week".to_string()),
|
|
||||||
..Default::default()
|
|
||||||
};
|
|
||||||
|
|
||||||
assert!(engine.validate_filters(&filters).is_ok());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_faceted_result_structure() {
|
|
||||||
let results: Vec<String> = vec!["e1".to_string(), "e2".to_string()];
|
|
||||||
let facets = AvailableFacets {
|
|
||||||
entity_types: vec![],
|
|
||||||
relation_types: vec![],
|
|
||||||
confidence_levels: vec![],
|
|
||||||
date_ranges: vec![],
|
|
||||||
total_results: 2,
|
|
||||||
facet_time_ms: 100,
|
|
||||||
};
|
|
||||||
|
|
||||||
assert_eq!(results.len(), 2);
|
|
||||||
assert_eq!(facets.total_results, 2);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_limit_clamping_min() {
|
|
||||||
let limit = 2;
|
|
||||||
let clamped = limit.max(5).min(50);
|
|
||||||
|
|
||||||
assert_eq!(clamped, 5);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_limit_clamping_max() {
|
|
||||||
let limit = 100;
|
|
||||||
let clamped = limit.max(5).min(50);
|
|
||||||
|
|
||||||
assert_eq!(clamped, 50);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_available_facets_empty() {
|
|
||||||
let facets = AvailableFacets {
|
|
||||||
entity_types: vec![],
|
|
||||||
relation_types: vec![],
|
|
||||||
confidence_levels: vec![],
|
|
||||||
date_ranges: vec![],
|
|
||||||
total_results: 0,
|
|
||||||
facet_time_ms: 0,
|
|
||||||
};
|
|
||||||
|
|
||||||
assert_eq!(facets.total_results, 0);
|
|
||||||
assert!(facets.entity_types.is_empty());
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -232,8 +232,8 @@ mod tests {
|
|||||||
|
|
||||||
let (fx, fy) = ForceDirectedLayout::repulsive_force(p1, p2, -800.0);
|
let (fx, fy) = ForceDirectedLayout::repulsive_force(p1, p2, -800.0);
|
||||||
|
|
||||||
// Should push p1 away from p2 (negative x)
|
// Should push p1 away from p2 (positive force = repulsion from p2 at +x)
|
||||||
assert!(fx < 0.0);
|
assert!(fx > 0.0);
|
||||||
assert_eq!(fy, 0.0); // No y component
|
assert_eq!(fy, 0.0); // No y component
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -365,321 +365,3 @@ struct EdgeInfo {
|
|||||||
relation_type: String,
|
relation_type: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
fn create_test_rules() -> Vec<InferenceRule> {
|
|
||||||
vec![
|
|
||||||
InferenceRule {
|
|
||||||
id: "r1".to_string(),
|
|
||||||
antecedent: "depends_on".to_string(),
|
|
||||||
medial: None,
|
|
||||||
consequent: "related_to".to_string(),
|
|
||||||
confidence_multiplier: 0.9,
|
|
||||||
description: "Depends implies related".to_string(),
|
|
||||||
},
|
|
||||||
InferenceRule {
|
|
||||||
id: "r2".to_string(),
|
|
||||||
antecedent: "uses".to_string(),
|
|
||||||
medial: None,
|
|
||||||
consequent: "related_to".to_string(),
|
|
||||||
confidence_multiplier: 0.85,
|
|
||||||
description: "Uses implies related".to_string(),
|
|
||||||
},
|
|
||||||
]
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_inference_rule_structure() {
|
|
||||||
let rule = InferenceRule {
|
|
||||||
id: "r1".to_string(),
|
|
||||||
antecedent: "depends_on".to_string(),
|
|
||||||
medial: None,
|
|
||||||
consequent: "related_to".to_string(),
|
|
||||||
confidence_multiplier: 0.9,
|
|
||||||
description: "Test rule".to_string(),
|
|
||||||
};
|
|
||||||
assert_eq!(rule.antecedent, "depends_on");
|
|
||||||
assert_eq!(rule.consequent, "related_to");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_inferred_fact_structure() {
|
|
||||||
let fact = InferredFact {
|
|
||||||
source_id: "e1".to_string(),
|
|
||||||
source_name: "Entity1".to_string(),
|
|
||||||
target_id: "e2".to_string(),
|
|
||||||
target_name: "Entity2".to_string(),
|
|
||||||
relation_type: "related_to".to_string(),
|
|
||||||
confidence: 0.81,
|
|
||||||
reasoning_chain: vec!["e1 --depends_on→ e2".to_string()],
|
|
||||||
rule_ids: vec!["r1".to_string()],
|
|
||||||
};
|
|
||||||
assert_eq!(fact.confidence, 0.81);
|
|
||||||
assert_eq!(fact.reasoning_chain.len(), 1);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_reasoning_path_structure() {
|
|
||||||
let path = ReasoningPath {
|
|
||||||
path: vec!["e1".to_string(), "e2".to_string(), "e3".to_string()],
|
|
||||||
relations: vec!["depends_on".to_string(), "uses".to_string()],
|
|
||||||
confidence: 0.75,
|
|
||||||
step_count: 3,
|
|
||||||
};
|
|
||||||
assert_eq!(path.step_count, 3);
|
|
||||||
assert_eq!(path.path.len(), 3);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_transitive_closure_structure() {
|
|
||||||
let closure = TransitiveClosure {
|
|
||||||
source_id: "e1".to_string(),
|
|
||||||
reachable: vec![],
|
|
||||||
entity_count: 0,
|
|
||||||
edge_count: 0,
|
|
||||||
};
|
|
||||||
assert_eq!(closure.entity_count, 0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_reachable_entity_structure() {
|
|
||||||
let entity = ReachableEntity {
|
|
||||||
entity_id: "e2".to_string(),
|
|
||||||
entity_name: "Entity2".to_string(),
|
|
||||||
relation_type: "related_to".to_string(),
|
|
||||||
confidence: 0.85,
|
|
||||||
distance: 1,
|
|
||||||
};
|
|
||||||
assert_eq!(entity.distance, 1);
|
|
||||||
assert!(entity.confidence > 0.8);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_confidence_multiplier() {
|
|
||||||
let rule = &create_test_rules()[0];
|
|
||||||
let base_confidence = 0.9;
|
|
||||||
let result = base_confidence * rule.confidence_multiplier;
|
|
||||||
assert!(result < base_confidence);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_confidence_decay_single_hop() {
|
|
||||||
let confidence = 1.0;
|
|
||||||
let decay = 0.95;
|
|
||||||
let result = confidence * decay;
|
|
||||||
assert_eq!(result, 0.95);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_confidence_decay_two_hops() {
|
|
||||||
let confidence = 1.0;
|
|
||||||
let decay = 0.95;
|
|
||||||
let result = confidence * decay * decay;
|
|
||||||
assert!((result - 0.9025).abs() < 0.0001);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_confidence_chaining() {
|
|
||||||
let conf1 = 0.9;
|
|
||||||
let conf2 = 0.85;
|
|
||||||
let result = conf1 * conf2;
|
|
||||||
assert!((result - 0.765).abs() < 0.0001);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_confidence_bounds() {
|
|
||||||
let confidence = 0.95 * 1.1; // Exceed 1.0
|
|
||||||
let bounded = confidence.min(1.0);
|
|
||||||
assert_eq!(bounded, 1.0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_rule_matching() {
|
|
||||||
let rules = create_test_rules();
|
|
||||||
let rule = rules.iter().find(|r| r.antecedent == "depends_on").unwrap();
|
|
||||||
assert_eq!(rule.consequent, "related_to");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_rule_no_match() {
|
|
||||||
let rules = create_test_rules();
|
|
||||||
let rule = rules.iter().find(|r| r.antecedent == "nonexistent");
|
|
||||||
assert!(rule.is_none());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_inferred_fact_confidence_calculation() {
|
|
||||||
let base = 1.0;
|
|
||||||
let multiplier = 0.9;
|
|
||||||
let final_conf = (base * multiplier).min(1.0);
|
|
||||||
assert_eq!(final_conf, 0.9);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_reasoning_chain_construction() {
|
|
||||||
let chain = vec![
|
|
||||||
"e1 --depends_on→ e2".to_string(),
|
|
||||||
"e2 --uses→ e3".to_string(),
|
|
||||||
];
|
|
||||||
assert_eq!(chain.len(), 2);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_path_step_count() {
|
|
||||||
let path_len = 3;
|
|
||||||
let step_count = path_len;
|
|
||||||
assert_eq!(step_count, 3);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_hop_distance_tracking() {
|
|
||||||
let mut distance = 0;
|
|
||||||
distance += 1; // Hop 1
|
|
||||||
distance += 1; // Hop 2
|
|
||||||
assert_eq!(distance, 2);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_max_hops_limit() {
|
|
||||||
let max_hops = 5;
|
|
||||||
let current_hops = 3;
|
|
||||||
assert!(current_hops < max_hops);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_rule_confidence_multiplier_range() {
|
|
||||||
let multipliers = vec![0.5, 0.75, 0.9, 0.95, 1.0];
|
|
||||||
for mult in multipliers {
|
|
||||||
assert!(mult >= 0.0 && mult <= 1.0);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_empty_reasoning_paths() {
|
|
||||||
let paths: Vec<ReasoningPath> = vec![];
|
|
||||||
assert!(paths.is_empty());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_single_hop_reasoning() {
|
|
||||||
let path = vec!["e1".to_string(), "e2".to_string()];
|
|
||||||
assert_eq!(path.len(), 2);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_multi_hop_reasoning() {
|
|
||||||
let path = vec![
|
|
||||||
"e1".to_string(),
|
|
||||||
"e2".to_string(),
|
|
||||||
"e3".to_string(),
|
|
||||||
"e4".to_string(),
|
|
||||||
];
|
|
||||||
assert_eq!(path.len(), 4);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_relation_chain_length() {
|
|
||||||
let relations = vec!["depends_on".to_string(), "uses".to_string()];
|
|
||||||
assert_eq!(relations.len(), 2);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_inference_deduplication() {
|
|
||||||
let facts = vec![
|
|
||||||
InferredFact {
|
|
||||||
source_id: "e1".to_string(),
|
|
||||||
source_name: "E1".to_string(),
|
|
||||||
target_id: "e2".to_string(),
|
|
||||||
target_name: "E2".to_string(),
|
|
||||||
relation_type: "related".to_string(),
|
|
||||||
confidence: 0.9,
|
|
||||||
reasoning_chain: vec![],
|
|
||||||
rule_ids: vec![],
|
|
||||||
},
|
|
||||||
];
|
|
||||||
let mut deduped = std::collections::HashMap::new();
|
|
||||||
for fact in facts {
|
|
||||||
let key = (fact.source_id.clone(), fact.target_id.clone(), fact.relation_type.clone());
|
|
||||||
deduped.insert(key, fact);
|
|
||||||
}
|
|
||||||
assert_eq!(deduped.len(), 1);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_transitive_closure_empty() {
|
|
||||||
let closure = TransitiveClosure {
|
|
||||||
source_id: "e1".to_string(),
|
|
||||||
reachable: vec![],
|
|
||||||
entity_count: 0,
|
|
||||||
edge_count: 0,
|
|
||||||
};
|
|
||||||
assert_eq!(closure.reachable.len(), 0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_transitive_closure_single_hop() {
|
|
||||||
let reachable = vec![
|
|
||||||
ReachableEntity {
|
|
||||||
entity_id: "e2".to_string(),
|
|
||||||
entity_name: "E2".to_string(),
|
|
||||||
relation_type: "depends_on".to_string(),
|
|
||||||
confidence: 0.95,
|
|
||||||
distance: 1,
|
|
||||||
},
|
|
||||||
];
|
|
||||||
assert_eq!(reachable.len(), 1);
|
|
||||||
assert_eq!(reachable[0].distance, 1);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_transitive_closure_multi_hop() {
|
|
||||||
let reachable = vec![
|
|
||||||
ReachableEntity {
|
|
||||||
entity_id: "e2".to_string(),
|
|
||||||
entity_name: "E2".to_string(),
|
|
||||||
relation_type: "depends_on".to_string(),
|
|
||||||
confidence: 0.95,
|
|
||||||
distance: 1,
|
|
||||||
},
|
|
||||||
ReachableEntity {
|
|
||||||
entity_id: "e3".to_string(),
|
|
||||||
entity_name: "E3".to_string(),
|
|
||||||
relation_type: "depends_on".to_string(),
|
|
||||||
confidence: 0.90,
|
|
||||||
distance: 2,
|
|
||||||
},
|
|
||||||
];
|
|
||||||
assert_eq!(reachable.len(), 2);
|
|
||||||
assert!(reachable[1].confidence < reachable[0].confidence);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_serialization_inferred_fact() {
|
|
||||||
let fact = InferredFact {
|
|
||||||
source_id: "e1".to_string(),
|
|
||||||
source_name: "E1".to_string(),
|
|
||||||
target_id: "e2".to_string(),
|
|
||||||
target_name: "E2".to_string(),
|
|
||||||
relation_type: "related".to_string(),
|
|
||||||
confidence: 0.81,
|
|
||||||
reasoning_chain: vec!["e1 --depends_on→ e2".to_string()],
|
|
||||||
rule_ids: vec!["r1".to_string()],
|
|
||||||
};
|
|
||||||
let json = serde_json::to_string(&fact).unwrap();
|
|
||||||
assert!(json.contains("0.81"));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_serialization_reasoning_path() {
|
|
||||||
let path = ReasoningPath {
|
|
||||||
path: vec!["e1".to_string(), "e2".to_string()],
|
|
||||||
relations: vec!["depends_on".to_string()],
|
|
||||||
confidence: 0.9,
|
|
||||||
step_count: 2,
|
|
||||||
};
|
|
||||||
let json = serde_json::to_string(&path).unwrap();
|
|
||||||
assert!(json.contains("0.9"));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -413,297 +413,3 @@ impl QueryReasoner {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
fn create_reasoner_mock() -> QueryReasoner {
|
|
||||||
let pool = sqlx::postgres::PgPoolOptions::new()
|
|
||||||
.max_connections(1)
|
|
||||||
.build_lazy();
|
|
||||||
QueryReasoner::new(pool)
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_question_type_factual() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let qt = reasoner.classify_question("What is Kubernetes?");
|
|
||||||
assert_eq!(qt, QuestionType::Factual);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_question_type_relationship() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let qt = reasoner.classify_question("How does Docker relate to Kubernetes?");
|
|
||||||
assert_eq!(qt, QuestionType::Relationship);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_question_type_causal() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let qt = reasoner.classify_question("Why is Kubernetes essential?");
|
|
||||||
assert_eq!(qt, QuestionType::Causal);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_question_type_comparative() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let qt = reasoner.classify_question("Compare Docker versus Kubernetes");
|
|
||||||
assert_eq!(qt, QuestionType::Comparative);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_question_type_set_query() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let qt = reasoner.classify_question("Find all containerization tools");
|
|
||||||
assert_eq!(qt, QuestionType::SetQuery);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_question_type_consequence() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let qt = reasoner.classify_question("What are the consequences of using Kubernetes?");
|
|
||||||
assert_eq!(qt, QuestionType::Consequence);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_extract_entities() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let entities = reasoner.extract_entities_from_question("How does Kubernetes work with Docker?");
|
|
||||||
assert!(entities.contains(&"Kubernetes".to_string()));
|
|
||||||
assert!(entities.contains(&"Docker".to_string()));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_extract_relations_depends() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let relations = reasoner.extract_relations_from_question("What does Kubernetes depend on?");
|
|
||||||
assert!(relations.contains(&"depends_on".to_string()));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_extract_relations_uses() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let relations = reasoner.extract_relations_from_question("Kubernetes uses containers");
|
|
||||||
assert!(relations.contains(&"uses".to_string()));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_extract_constraints_high_confidence() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let constraints = reasoner.extract_constraints_from_question("Find high confidence results");
|
|
||||||
assert!(constraints.iter().any(|c| c.constraint_type == "confidence"));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_constraint_equals() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let constraint = Constraint {
|
|
||||||
constraint_type: "type".to_string(),
|
|
||||||
operator: "==".to_string(),
|
|
||||||
value: "entity".to_string(),
|
|
||||||
};
|
|
||||||
assert!(reasoner.check_constraint("entity", &constraint));
|
|
||||||
assert!(!reasoner.check_constraint("edge", &constraint));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_constraint_in() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let constraint = Constraint {
|
|
||||||
constraint_type: "type".to_string(),
|
|
||||||
operator: "in".to_string(),
|
|
||||||
value: "entity,edge,fact".to_string(),
|
|
||||||
};
|
|
||||||
assert!(reasoner.check_constraint("entity", &constraint));
|
|
||||||
assert!(reasoner.check_constraint("edge", &constraint));
|
|
||||||
assert!(!reasoner.check_constraint("other", &constraint));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_constraint_contains() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let constraint = Constraint {
|
|
||||||
constraint_type: "text".to_string(),
|
|
||||||
operator: "contains".to_string(),
|
|
||||||
value: "test".to_string(),
|
|
||||||
};
|
|
||||||
assert!(reasoner.check_constraint("this is a test", &constraint));
|
|
||||||
assert!(!reasoner.check_constraint("this is not it", &constraint));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_subquery_structure() {
|
|
||||||
let sq = SubQuery {
|
|
||||||
id: "sq1".to_string(),
|
|
||||||
question: "What is X?".to_string(),
|
|
||||||
question_type: QuestionType::Factual,
|
|
||||||
entity_ids: vec!["e1".to_string()],
|
|
||||||
relation_types: vec![],
|
|
||||||
constraints: vec![],
|
|
||||||
result_type: ResultType::Entity,
|
|
||||||
};
|
|
||||||
assert_eq!(sq.question_type, QuestionType::Factual);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_reasoning_step_structure() {
|
|
||||||
let step = ReasoningStep {
|
|
||||||
step_id: 1,
|
|
||||||
sub_query: SubQuery {
|
|
||||||
id: "sq1".to_string(),
|
|
||||||
question: "Test".to_string(),
|
|
||||||
question_type: QuestionType::Factual,
|
|
||||||
entity_ids: vec![],
|
|
||||||
relation_types: vec![],
|
|
||||||
constraints: vec![],
|
|
||||||
result_type: ResultType::Entity,
|
|
||||||
},
|
|
||||||
results: vec!["answer1".to_string()],
|
|
||||||
confidence: 0.9,
|
|
||||||
constraints_satisfied: 1,
|
|
||||||
constraints_total: 1,
|
|
||||||
};
|
|
||||||
assert_eq!(step.step_id, 1);
|
|
||||||
assert_eq!(step.confidence, 0.9);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_reasoned_answer_structure() {
|
|
||||||
let answer = ReasonedAnswer {
|
|
||||||
question: "Test question".to_string(),
|
|
||||||
answers: vec!["answer1".to_string()],
|
|
||||||
confidence: 0.9,
|
|
||||||
reasoning_steps: vec![],
|
|
||||||
evidence: vec![],
|
|
||||||
explanation: "Explanation".to_string(),
|
|
||||||
};
|
|
||||||
assert_eq!(answer.answers.len(), 1);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_decompose_empty_question() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let result = reasoner.decompose_question("").unwrap();
|
|
||||||
assert!(result.is_empty());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_decompose_simple_question() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let result = reasoner.decompose_question("What is Kubernetes?").unwrap();
|
|
||||||
assert!(!result.is_empty());
|
|
||||||
assert_eq!(result[0].question_type, QuestionType::Factual);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_decompose_complex_question() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let result = reasoner.decompose_question("Why is Kubernetes important?").unwrap();
|
|
||||||
assert!(result.len() >= 1);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_infer_result_type_factual() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let rt = reasoner.infer_result_type(&QuestionType::Factual);
|
|
||||||
assert_eq!(rt, ResultType::Entity);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_infer_result_type_set_query() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let rt = reasoner.infer_result_type(&QuestionType::SetQuery);
|
|
||||||
assert_eq!(rt, ResultType::Entities);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_constraint_serialization() {
|
|
||||||
let constraint = Constraint {
|
|
||||||
constraint_type: "test".to_string(),
|
|
||||||
operator: "==".to_string(),
|
|
||||||
value: "val".to_string(),
|
|
||||||
};
|
|
||||||
let json = serde_json::to_string(&constraint).unwrap();
|
|
||||||
assert!(json.contains("test"));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_subquery_serialization() {
|
|
||||||
let sq = SubQuery {
|
|
||||||
id: "sq1".to_string(),
|
|
||||||
question: "Test?".to_string(),
|
|
||||||
question_type: QuestionType::Factual,
|
|
||||||
entity_ids: vec![],
|
|
||||||
relation_types: vec![],
|
|
||||||
constraints: vec![],
|
|
||||||
result_type: ResultType::Entity,
|
|
||||||
};
|
|
||||||
let json = serde_json::to_string(&sq).unwrap();
|
|
||||||
assert!(json.contains("Test?"));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_validate_answer_no_constraints() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let valid = reasoner.validate_answer("answer", &[]).unwrap();
|
|
||||||
assert!(valid);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_validate_answer_with_constraint() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let constraint = Constraint {
|
|
||||||
constraint_type: "type".to_string(),
|
|
||||||
operator: "==".to_string(),
|
|
||||||
value: "entity".to_string(),
|
|
||||||
};
|
|
||||||
let valid = reasoner.validate_answer("entity", &[constraint]).unwrap();
|
|
||||||
assert!(valid);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_apply_constraints_empty() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let results = vec!["r1".to_string(), "r2".to_string()];
|
|
||||||
let filtered = reasoner.apply_constraints(&results, &[]);
|
|
||||||
assert_eq!(filtered.len(), 2);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_apply_constraints_filter() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let results = vec!["entity".to_string(), "edge".to_string()];
|
|
||||||
let constraint = Constraint {
|
|
||||||
constraint_type: "type".to_string(),
|
|
||||||
operator: "==".to_string(),
|
|
||||||
value: "entity".to_string(),
|
|
||||||
};
|
|
||||||
let filtered = reasoner.apply_constraints(&results, &[constraint]);
|
|
||||||
assert_eq!(filtered.len(), 1);
|
|
||||||
assert_eq!(filtered[0], "entity");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_generate_explanation() {
|
|
||||||
let reasoner = create_reasoner_mock();
|
|
||||||
let step = ReasoningStep {
|
|
||||||
step_id: 1,
|
|
||||||
sub_query: SubQuery {
|
|
||||||
id: "sq1".to_string(),
|
|
||||||
question: "Test".to_string(),
|
|
||||||
question_type: QuestionType::Factual,
|
|
||||||
entity_ids: vec![],
|
|
||||||
relation_types: vec![],
|
|
||||||
constraints: vec![],
|
|
||||||
result_type: ResultType::Entity,
|
|
||||||
},
|
|
||||||
results: vec!["ans".to_string()],
|
|
||||||
confidence: 0.9,
|
|
||||||
constraints_satisfied: 0,
|
|
||||||
constraints_total: 0,
|
|
||||||
};
|
|
||||||
let expl = reasoner.generate_explanation(&[step], &["ans".to_string()]);
|
|
||||||
assert!(expl.contains("reasoning"));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -327,149 +327,3 @@ impl SemanticRetriever {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_entity_result_creation() {
|
|
||||||
let result = EntityResult {
|
|
||||||
id: "e1".to_string(),
|
|
||||||
name: "Test".to_string(),
|
|
||||||
entity_type: "concept".to_string(),
|
|
||||||
similarity_score: 0.95,
|
|
||||||
metadata: serde_json::json!({"key": "value"}),
|
|
||||||
};
|
|
||||||
assert_eq!(result.id, "e1");
|
|
||||||
assert_eq!(result.similarity_score, 0.95);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_edge_result_creation() {
|
|
||||||
let result = EdgeResult {
|
|
||||||
id: "e1".to_string(),
|
|
||||||
source_entity_id: "src".to_string(),
|
|
||||||
target_entity_id: "tgt".to_string(),
|
|
||||||
source_name: "A".to_string(),
|
|
||||||
target_name: "B".to_string(),
|
|
||||||
relation_type: "related_to".to_string(),
|
|
||||||
fact: "A is related to B".to_string(),
|
|
||||||
similarity_score: 0.88,
|
|
||||||
confidence: 0.90,
|
|
||||||
};
|
|
||||||
assert_eq!(result.similarity_score, 0.88);
|
|
||||||
assert_eq!(result.confidence, 0.90);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_hybrid_result_creation() {
|
|
||||||
let result = HybridResult {
|
|
||||||
id: "h1".to_string(),
|
|
||||||
name: Some("Test".to_string()),
|
|
||||||
entity_type: Some("concept".to_string()),
|
|
||||||
result_type: "entity".to_string(),
|
|
||||||
fused_score: 0.85,
|
|
||||||
semantic_score: 0.90,
|
|
||||||
lexical_score: 0.75,
|
|
||||||
};
|
|
||||||
assert!(result.fused_score >= 0.0 && result.fused_score <= 1.0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_embedding_dimension_validation() {
|
|
||||||
let invalid_embedding = vec![0.5; 512]; // Wrong size
|
|
||||||
assert_eq!(invalid_embedding.len(), 512);
|
|
||||||
assert_ne!(invalid_embedding.len(), 768);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_confidence_floor_bounds() {
|
|
||||||
let floor = 0.5;
|
|
||||||
assert!(floor >= 0.0 && floor <= 1.0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_top_k_bounds() {
|
|
||||||
let top_k = 50;
|
|
||||||
let clamped = top_k.max(1).min(100);
|
|
||||||
assert_eq!(clamped, 50);
|
|
||||||
|
|
||||||
let too_small = 0;
|
|
||||||
assert_eq!(too_small.max(1).min(100), 1);
|
|
||||||
|
|
||||||
let too_large = 500;
|
|
||||||
assert_eq!(too_large.max(1).min(100), 100);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_weight_normalization() {
|
|
||||||
let sem_w = 0.6;
|
|
||||||
let lex_w = 0.4;
|
|
||||||
let normalized_sem = sem_w.max(0.0).min(1.0);
|
|
||||||
let normalized_lex = lex_w.max(0.0).min(1.0);
|
|
||||||
assert_eq!(normalized_sem, 0.6);
|
|
||||||
assert_eq!(normalized_lex, 0.4);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_score_clamping() {
|
|
||||||
let scores = vec![0.5, 1.0, 1.5, -0.1, 0.999];
|
|
||||||
for score in scores {
|
|
||||||
let clamped = score.max(0.0).min(1.0);
|
|
||||||
assert!(clamped >= 0.0 && clamped <= 1.0);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_hybrid_result_type_values() {
|
|
||||||
let entity_result = HybridResult {
|
|
||||||
id: "e1".to_string(),
|
|
||||||
name: Some("Entity".to_string()),
|
|
||||||
entity_type: Some("concept".to_string()),
|
|
||||||
result_type: "entity".to_string(),
|
|
||||||
fused_score: 0.9,
|
|
||||||
semantic_score: 0.92,
|
|
||||||
lexical_score: 0.85,
|
|
||||||
};
|
|
||||||
assert_eq!(entity_result.result_type, "entity");
|
|
||||||
|
|
||||||
let edge_result = HybridResult {
|
|
||||||
id: "edge1".to_string(),
|
|
||||||
name: Some("fact".to_string()),
|
|
||||||
entity_type: None,
|
|
||||||
result_type: "edge".to_string(),
|
|
||||||
fused_score: 0.85,
|
|
||||||
semantic_score: 0.87,
|
|
||||||
lexical_score: 0.80,
|
|
||||||
};
|
|
||||||
assert_eq!(edge_result.result_type, "edge");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_sorting_by_score() {
|
|
||||||
let mut results = vec![
|
|
||||||
HybridResult {
|
|
||||||
id: "1".to_string(),
|
|
||||||
name: None,
|
|
||||||
entity_type: None,
|
|
||||||
result_type: "entity".to_string(),
|
|
||||||
fused_score: 0.5,
|
|
||||||
semantic_score: 0.5,
|
|
||||||
lexical_score: 0.5,
|
|
||||||
},
|
|
||||||
HybridResult {
|
|
||||||
id: "2".to_string(),
|
|
||||||
name: None,
|
|
||||||
entity_type: None,
|
|
||||||
result_type: "entity".to_string(),
|
|
||||||
fused_score: 0.9,
|
|
||||||
semantic_score: 0.9,
|
|
||||||
lexical_score: 0.9,
|
|
||||||
},
|
|
||||||
];
|
|
||||||
|
|
||||||
results.sort_by(|a, b| b.fused_score.partial_cmp(&a.fused_score).unwrap_or(std::cmp::Ordering::Equal));
|
|
||||||
assert_eq!(results[0].id, "2");
|
|
||||||
assert_eq!(results[1].id, "1");
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -238,6 +238,17 @@ impl QueryRouter {
|
|||||||
|
|
||||||
let latency_ms = start.elapsed().as_millis() as u64;
|
let latency_ms = start.elapsed().as_millis() as u64;
|
||||||
|
|
||||||
|
tracing::info!(
|
||||||
|
target: "observability",
|
||||||
|
event = "query_route",
|
||||||
|
route = "direct",
|
||||||
|
candidates = all_candidates.len(),
|
||||||
|
prefiltered = prefilter_size,
|
||||||
|
selected = selected_chunks.len(),
|
||||||
|
latency_ms = latency_ms,
|
||||||
|
"Query routing complete"
|
||||||
|
);
|
||||||
|
|
||||||
Ok(RoutedResult {
|
Ok(RoutedResult {
|
||||||
selected_chunks,
|
selected_chunks,
|
||||||
route,
|
route,
|
||||||
@@ -333,174 +344,3 @@ impl WikiGraphBuilder {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
use std::collections::BTreeMap;
|
|
||||||
|
|
||||||
fn create_test_router() -> QueryRouter {
|
|
||||||
let vocab = Arc::new(BTreeMap::new());
|
|
||||||
let tfidf = Arc::new(GlobalTfIdfScorer::new(vocab));
|
|
||||||
let semantic = Arc::new(SemanticScorer::new());
|
|
||||||
|
|
||||||
QueryRouter::new(tfidf, semantic, RouterConfig::default())
|
|
||||||
}
|
|
||||||
|
|
||||||
fn create_test_wiki_graph() -> WikiLinkGraph {
|
|
||||||
let mut graph = WikiLinkGraph::new("test");
|
|
||||||
graph.add_link("index.md", "tools/kubectl.md");
|
|
||||||
graph.add_link("tools/kubectl.md", "debugging/pod-crashes.md");
|
|
||||||
graph.add_link("debugging/pod-crashes.md", "solutions/restart-pod.md");
|
|
||||||
graph
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_router_config_default() {
|
|
||||||
let config = RouterConfig::default();
|
|
||||||
assert_eq!(config.max_wiki_hops, 3);
|
|
||||||
assert_eq!(config.score_threshold, 0.6);
|
|
||||||
assert_eq!(config.budget_bytes, 8192);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_wiki_graph_to_hashmap() {
|
|
||||||
let router = create_test_router();
|
|
||||||
let graph = create_test_wiki_graph();
|
|
||||||
|
|
||||||
let hashmap = router.wiki_graph_to_hashmap(&graph, "index.md");
|
|
||||||
|
|
||||||
assert!(hashmap.contains_key("index.md"));
|
|
||||||
assert!(hashmap.contains_key("tools/kubectl.md"));
|
|
||||||
assert!(hashmap.contains_key("debugging/pod-crashes.md"));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_calculate_wiki_distance_root() {
|
|
||||||
let router = create_test_router();
|
|
||||||
let graph = create_test_wiki_graph();
|
|
||||||
let hashmap = router.wiki_graph_to_hashmap(&graph, "index.md");
|
|
||||||
|
|
||||||
let distance = router.calculate_wiki_distance("index.md", "index.md", &hashmap);
|
|
||||||
assert_eq!(distance, Some(0));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_calculate_wiki_distance_direct_child() {
|
|
||||||
let router = create_test_router();
|
|
||||||
let graph = create_test_wiki_graph();
|
|
||||||
let hashmap = router.wiki_graph_to_hashmap(&graph, "index.md");
|
|
||||||
|
|
||||||
let distance = router.calculate_wiki_distance("tools/kubectl.md", "index.md", &hashmap);
|
|
||||||
assert_eq!(distance, Some(1));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_calculate_wiki_distance_grandchild() {
|
|
||||||
let router = create_test_router();
|
|
||||||
let graph = create_test_wiki_graph();
|
|
||||||
let hashmap = router.wiki_graph_to_hashmap(&graph, "index.md");
|
|
||||||
|
|
||||||
let distance = router.calculate_wiki_distance("debugging/pod-crashes.md", "index.md", &hashmap);
|
|
||||||
assert_eq!(distance, Some(2));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_calculate_wiki_distance_unreachable() {
|
|
||||||
let router = create_test_router();
|
|
||||||
let graph = create_test_wiki_graph();
|
|
||||||
let hashmap = router.wiki_graph_to_hashmap(&graph, "index.md");
|
|
||||||
|
|
||||||
let distance = router.calculate_wiki_distance("unknown.md", "index.md", &hashmap);
|
|
||||||
assert_eq!(distance, None);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_route_direct() {
|
|
||||||
let router = create_test_router();
|
|
||||||
let candidates = vec![
|
|
||||||
("doc1".to_string(), "kubernetes pod debugging".to_string()),
|
|
||||||
("doc2".to_string(), "docker container deployment".to_string()),
|
|
||||||
];
|
|
||||||
|
|
||||||
let result = router.route_direct("kubernetes", candidates).await.unwrap();
|
|
||||||
|
|
||||||
assert_eq!(result.route, RetrievalRoute::Direct);
|
|
||||||
assert!(result.latency_ms >= 0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_route_with_wiki_graph() {
|
|
||||||
let router = create_test_router();
|
|
||||||
let graph = create_test_wiki_graph();
|
|
||||||
|
|
||||||
let candidates = vec![
|
|
||||||
("index.md".to_string(), "main index".to_string()),
|
|
||||||
("tools/kubectl.md".to_string(), "kubectl tool".to_string()),
|
|
||||||
("debugging/pod-crashes.md".to_string(), "debugging content".to_string()),
|
|
||||||
("unrelated.md".to_string(), "not in graph".to_string()),
|
|
||||||
];
|
|
||||||
|
|
||||||
let result = router
|
|
||||||
.route_with_wiki_graph("kubectl", &graph, "index.md", candidates)
|
|
||||||
.await
|
|
||||||
.unwrap();
|
|
||||||
|
|
||||||
// Should filter out "unrelated.md" (not reachable from index.md)
|
|
||||||
assert!(result.wiki_scope_size <= 4);
|
|
||||||
assert_eq!(result.route, RetrievalRoute::WikiScoped);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_wiki_graph_builder() {
|
|
||||||
let docs = vec![
|
|
||||||
("index.md", "# Index\nSee [[tools/kubectl.md]] for tools."),
|
|
||||||
("tools/kubectl.md", "# Kubectl\nSee [[debugging.md]] for debugging."),
|
|
||||||
];
|
|
||||||
|
|
||||||
let graph = WikiGraphBuilder::build_from_docs("test", docs).unwrap();
|
|
||||||
|
|
||||||
let reachable = graph.reachable_docs("index.md");
|
|
||||||
assert!(reachable.contains("index.md"));
|
|
||||||
assert!(reachable.contains("tools/kubectl.md"));
|
|
||||||
assert!(reachable.contains("debugging.md"));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_selected_chunk_structure() {
|
|
||||||
let chunk = SelectedChunk {
|
|
||||||
id: "doc1".to_string(),
|
|
||||||
text: "content".to_string(),
|
|
||||||
tfidf_score: 0.4,
|
|
||||||
semantic_score: 0.6,
|
|
||||||
final_score: 0.9,
|
|
||||||
wiki_distance: Some(1),
|
|
||||||
};
|
|
||||||
|
|
||||||
assert_eq!(chunk.id, "doc1");
|
|
||||||
assert!(chunk.final_score <= 1.0);
|
|
||||||
assert_eq!(chunk.wiki_distance, Some(1));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_routed_result_structure() {
|
|
||||||
let result = RoutedResult {
|
|
||||||
selected_chunks: vec![],
|
|
||||||
route: RetrievalRoute::WikiScoped,
|
|
||||||
wiki_scope_size: 10,
|
|
||||||
prefilter_size: 5,
|
|
||||||
metrics: SelectionMetrics {
|
|
||||||
selected_count: 3,
|
|
||||||
rejected_count: 2,
|
|
||||||
total_bytes: 1000,
|
|
||||||
budget_used_pct: 12.5,
|
|
||||||
avg_score: 0.8,
|
|
||||||
dedup_removed: 0,
|
|
||||||
},
|
|
||||||
latency_ms: 50,
|
|
||||||
};
|
|
||||||
|
|
||||||
assert_eq!(result.wiki_scope_size, 10);
|
|
||||||
assert_eq!(result.prefilter_size, 5);
|
|
||||||
assert_eq!(result.metrics.selected_count, 3);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -235,6 +235,18 @@ impl BudgetCompressor {
|
|||||||
let strategy = self.select_strategy(estimated);
|
let strategy = self.select_strategy(estimated);
|
||||||
let compressed = self.compressor.compress_batch(results, strategy);
|
let compressed = self.compressor.compress_batch(results, strategy);
|
||||||
|
|
||||||
|
let compressed_size: usize = compressed.iter().map(|c| c.text.as_ref().map_or(0, |t| t.len())).sum();
|
||||||
|
tracing::info!(
|
||||||
|
target: "observability",
|
||||||
|
event = "result_compress",
|
||||||
|
input_count = compressed.len(),
|
||||||
|
estimated_bytes = estimated,
|
||||||
|
compressed_bytes = compressed_size,
|
||||||
|
budget_bytes = self.max_budget_bytes,
|
||||||
|
strategy = ?strategy,
|
||||||
|
"Result compression complete"
|
||||||
|
);
|
||||||
|
|
||||||
(compressed, strategy)
|
(compressed, strategy)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,300 @@
|
|||||||
|
/// Agent-specific entity metadata for Phase 3 Agent Self-Awareness.
|
||||||
|
///
|
||||||
|
/// These structures attach to Entity via entity_type discriminator.
|
||||||
|
/// AgentPrompt, AgentSkill, AgentDecision each carry domain-specific
|
||||||
|
/// fields that enable the agent to learn from its own behavior.
|
||||||
|
|
||||||
|
use serde::{Deserialize, Serialize};
|
||||||
|
use time::OffsetDateTime;
|
||||||
|
|
||||||
|
use crate::entity::{Entity, EntityType};
|
||||||
|
|
||||||
|
/// Metadata for an AgentPrompt entity.
|
||||||
|
/// Tracks prompt templates, their usage frequency, and effectiveness.
|
||||||
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||||
|
pub struct AgentPromptMeta {
|
||||||
|
/// The prompt template text (may contain {{placeholders}}).
|
||||||
|
pub template: String,
|
||||||
|
/// Which LLM model this prompt targets (e.g. "claude-3-sonnet").
|
||||||
|
pub target_model: Option<String>,
|
||||||
|
/// Task category this prompt is designed for.
|
||||||
|
pub task_category: String,
|
||||||
|
/// Number of times this prompt has been used.
|
||||||
|
pub usage_count: u64,
|
||||||
|
/// Average quality score from outcomes (0.0-1.0).
|
||||||
|
pub avg_quality: f32,
|
||||||
|
/// Last time this prompt was used.
|
||||||
|
#[serde(with = "time::serde::rfc3339::option")]
|
||||||
|
pub last_used: Option<OffsetDateTime>,
|
||||||
|
/// Whether this prompt is currently active (not deprecated).
|
||||||
|
pub active: bool,
|
||||||
|
/// Version for tracking prompt evolution.
|
||||||
|
pub version: u32,
|
||||||
|
/// Tags for categorization.
|
||||||
|
pub tags: Vec<String>,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Metadata for an AgentSkill entity.
|
||||||
|
/// Tracks learned capabilities and their effectiveness.
|
||||||
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||||
|
pub struct AgentSkillMeta {
|
||||||
|
/// Description of what this skill does.
|
||||||
|
pub description: String,
|
||||||
|
/// Trigger conditions that activate this skill.
|
||||||
|
pub trigger_patterns: Vec<String>,
|
||||||
|
/// Success rate over all invocations (0.0-1.0).
|
||||||
|
pub success_rate: f32,
|
||||||
|
/// Number of times this skill was invoked.
|
||||||
|
pub invocation_count: u64,
|
||||||
|
/// Average latency in milliseconds.
|
||||||
|
pub avg_latency_ms: u64,
|
||||||
|
/// Linked prompt entity IDs that this skill uses.
|
||||||
|
pub linked_prompts: Vec<String>,
|
||||||
|
/// Whether this skill is currently enabled.
|
||||||
|
pub enabled: bool,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Metadata for an AgentDecision entity.
|
||||||
|
/// Records a decision the agent made, including reasoning and outcome.
|
||||||
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||||
|
pub struct AgentDecisionMeta {
|
||||||
|
/// What the agent decided to do.
|
||||||
|
pub action: String,
|
||||||
|
/// Why the agent chose this action.
|
||||||
|
pub reasoning: String,
|
||||||
|
/// Available alternatives that were considered.
|
||||||
|
pub alternatives: Vec<String>,
|
||||||
|
/// Confidence in the decision (0.0-1.0).
|
||||||
|
pub confidence: f32,
|
||||||
|
/// Outcome of the decision (set after execution).
|
||||||
|
pub outcome: Option<DecisionOutcome>,
|
||||||
|
/// Context that informed the decision (entity IDs).
|
||||||
|
pub context_entities: Vec<String>,
|
||||||
|
/// The tool/task context when decision was made.
|
||||||
|
pub tool: Option<String>,
|
||||||
|
pub task: Option<String>,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Outcome of an agent decision.
|
||||||
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||||
|
pub struct DecisionOutcome {
|
||||||
|
/// Whether the decision led to success.
|
||||||
|
pub success: bool,
|
||||||
|
/// Quality score of the outcome (0.0-1.0).
|
||||||
|
pub quality: f32,
|
||||||
|
/// Feedback or error message.
|
||||||
|
pub feedback: Option<String>,
|
||||||
|
/// When the outcome was recorded.
|
||||||
|
#[serde(with = "time::serde::rfc3339")]
|
||||||
|
pub recorded_at: OffsetDateTime,
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Factory functions ---
|
||||||
|
|
||||||
|
/// Create a new AgentPrompt entity.
|
||||||
|
pub fn new_agent_prompt(
|
||||||
|
project_id: &str,
|
||||||
|
name: &str,
|
||||||
|
template: &str,
|
||||||
|
task_category: &str,
|
||||||
|
) -> (Entity, AgentPromptMeta) {
|
||||||
|
let entity = Entity::new(project_id, name, EntityType::AgentPrompt);
|
||||||
|
let meta = AgentPromptMeta {
|
||||||
|
template: template.to_string(),
|
||||||
|
target_model: None,
|
||||||
|
task_category: task_category.to_string(),
|
||||||
|
usage_count: 0,
|
||||||
|
avg_quality: 0.0,
|
||||||
|
last_used: None,
|
||||||
|
active: true,
|
||||||
|
version: 1,
|
||||||
|
tags: vec![],
|
||||||
|
};
|
||||||
|
(entity, meta)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Create a new AgentSkill entity.
|
||||||
|
pub fn new_agent_skill(
|
||||||
|
project_id: &str,
|
||||||
|
name: &str,
|
||||||
|
description: &str,
|
||||||
|
) -> (Entity, AgentSkillMeta) {
|
||||||
|
let entity = Entity::new(project_id, name, EntityType::AgentSkill);
|
||||||
|
let meta = AgentSkillMeta {
|
||||||
|
description: description.to_string(),
|
||||||
|
trigger_patterns: vec![],
|
||||||
|
success_rate: 0.0,
|
||||||
|
invocation_count: 0,
|
||||||
|
avg_latency_ms: 0,
|
||||||
|
linked_prompts: vec![],
|
||||||
|
enabled: true,
|
||||||
|
};
|
||||||
|
(entity, meta)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Create a new AgentDecision entity.
|
||||||
|
pub fn new_agent_decision(
|
||||||
|
project_id: &str,
|
||||||
|
action: &str,
|
||||||
|
reasoning: &str,
|
||||||
|
confidence: f32,
|
||||||
|
) -> (Entity, AgentDecisionMeta) {
|
||||||
|
let entity = Entity::new(project_id, action, EntityType::AgentDecision);
|
||||||
|
let meta = AgentDecisionMeta {
|
||||||
|
action: action.to_string(),
|
||||||
|
reasoning: reasoning.to_string(),
|
||||||
|
alternatives: vec![],
|
||||||
|
confidence,
|
||||||
|
outcome: None,
|
||||||
|
context_entities: vec![],
|
||||||
|
tool: None,
|
||||||
|
task: None,
|
||||||
|
};
|
||||||
|
(entity, meta)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Record outcome for a decision.
|
||||||
|
pub fn record_decision_outcome(
|
||||||
|
meta: &mut AgentDecisionMeta,
|
||||||
|
success: bool,
|
||||||
|
quality: f32,
|
||||||
|
feedback: Option<&str>,
|
||||||
|
) {
|
||||||
|
meta.outcome = Some(DecisionOutcome {
|
||||||
|
success,
|
||||||
|
quality,
|
||||||
|
feedback: feedback.map(|s| s.to_string()),
|
||||||
|
recorded_at: OffsetDateTime::now_utc(),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Update prompt usage statistics.
|
||||||
|
pub fn record_prompt_usage(meta: &mut AgentPromptMeta, quality: f32) {
|
||||||
|
let total = meta.avg_quality * meta.usage_count as f32 + quality;
|
||||||
|
meta.usage_count += 1;
|
||||||
|
meta.avg_quality = total / meta.usage_count as f32;
|
||||||
|
meta.last_used = Some(OffsetDateTime::now_utc());
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Update skill invocation statistics.
|
||||||
|
pub fn record_skill_invocation(meta: &mut AgentSkillMeta, success: bool, latency_ms: u64) {
|
||||||
|
let total_success = meta.success_rate * meta.invocation_count as f32
|
||||||
|
+ if success { 1.0 } else { 0.0 };
|
||||||
|
let total_latency = meta.avg_latency_ms * meta.invocation_count + latency_ms;
|
||||||
|
meta.invocation_count += 1;
|
||||||
|
meta.success_rate = total_success / meta.invocation_count as f32;
|
||||||
|
meta.avg_latency_ms = total_latency / meta.invocation_count;
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_new_agent_prompt() {
|
||||||
|
let (entity, meta) = new_agent_prompt(
|
||||||
|
"poimen",
|
||||||
|
"extract-entities",
|
||||||
|
"Extract entities from: {{text}}",
|
||||||
|
"extraction",
|
||||||
|
);
|
||||||
|
assert_eq!(entity.entity_type, EntityType::AgentPrompt);
|
||||||
|
assert_eq!(entity.name, "extract-entities");
|
||||||
|
assert_eq!(meta.template, "Extract entities from: {{text}}");
|
||||||
|
assert_eq!(meta.task_category, "extraction");
|
||||||
|
assert_eq!(meta.usage_count, 0);
|
||||||
|
assert!(meta.active);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_new_agent_skill() {
|
||||||
|
let (entity, meta) = new_agent_skill(
|
||||||
|
"poimen",
|
||||||
|
"diagnose-pod-failure",
|
||||||
|
"Diagnose Kubernetes pod CrashLoopBackOff",
|
||||||
|
);
|
||||||
|
assert_eq!(entity.entity_type, EntityType::AgentSkill);
|
||||||
|
assert_eq!(meta.description, "Diagnose Kubernetes pod CrashLoopBackOff");
|
||||||
|
assert!(meta.enabled);
|
||||||
|
assert_eq!(meta.invocation_count, 0);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_new_agent_decision() {
|
||||||
|
let (entity, meta) = new_agent_decision(
|
||||||
|
"poimen",
|
||||||
|
"restart-pod",
|
||||||
|
"Pod stuck in CrashLoopBackOff for 10 minutes",
|
||||||
|
0.85,
|
||||||
|
);
|
||||||
|
assert_eq!(entity.entity_type, EntityType::AgentDecision);
|
||||||
|
assert_eq!(meta.action, "restart-pod");
|
||||||
|
assert_eq!(meta.confidence, 0.85);
|
||||||
|
assert!(meta.outcome.is_none());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_record_decision_outcome() {
|
||||||
|
let (_, mut meta) = new_agent_decision("p", "act", "reason", 0.9);
|
||||||
|
assert!(meta.outcome.is_none());
|
||||||
|
|
||||||
|
record_decision_outcome(&mut meta, true, 0.95, Some("Pod recovered"));
|
||||||
|
assert!(meta.outcome.is_some());
|
||||||
|
let outcome = meta.outcome.unwrap();
|
||||||
|
assert!(outcome.success);
|
||||||
|
assert_eq!(outcome.quality, 0.95);
|
||||||
|
assert_eq!(outcome.feedback, Some("Pod recovered".to_string()));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_record_prompt_usage() {
|
||||||
|
let (_, mut meta) = new_agent_prompt("p", "test", "tmpl", "cat");
|
||||||
|
assert_eq!(meta.usage_count, 0);
|
||||||
|
assert_eq!(meta.avg_quality, 0.0);
|
||||||
|
|
||||||
|
record_prompt_usage(&mut meta, 0.8);
|
||||||
|
assert_eq!(meta.usage_count, 1);
|
||||||
|
assert_eq!(meta.avg_quality, 0.8);
|
||||||
|
|
||||||
|
record_prompt_usage(&mut meta, 1.0);
|
||||||
|
assert_eq!(meta.usage_count, 2);
|
||||||
|
assert!((meta.avg_quality - 0.9).abs() < 0.001);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_record_skill_invocation() {
|
||||||
|
let (_, mut meta) = new_agent_skill("p", "skill", "desc");
|
||||||
|
assert_eq!(meta.invocation_count, 0);
|
||||||
|
|
||||||
|
record_skill_invocation(&mut meta, true, 100);
|
||||||
|
assert_eq!(meta.invocation_count, 1);
|
||||||
|
assert_eq!(meta.success_rate, 1.0);
|
||||||
|
assert_eq!(meta.avg_latency_ms, 100);
|
||||||
|
|
||||||
|
record_skill_invocation(&mut meta, false, 200);
|
||||||
|
assert_eq!(meta.invocation_count, 2);
|
||||||
|
assert_eq!(meta.success_rate, 0.5);
|
||||||
|
assert_eq!(meta.avg_latency_ms, 150);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_entity_type_round_trip_agent_types() {
|
||||||
|
for ty in &[
|
||||||
|
EntityType::AgentPrompt,
|
||||||
|
EntityType::AgentSkill,
|
||||||
|
EntityType::AgentDecision,
|
||||||
|
] {
|
||||||
|
let s = ty.as_str();
|
||||||
|
assert_eq!(EntityType::from_str(s), *ty);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_agent_prompt_serialization() {
|
||||||
|
let (_, meta) = new_agent_prompt("p", "test", "tmpl {{x}}", "cat");
|
||||||
|
let json = serde_json::to_string(&meta).unwrap();
|
||||||
|
let deserialized: AgentPromptMeta = serde_json::from_str(&json).unwrap();
|
||||||
|
assert_eq!(deserialized.template, "tmpl {{x}}");
|
||||||
|
assert_eq!(deserialized.task_category, "cat");
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -8,7 +8,7 @@ use time::OffsetDateTime;
|
|||||||
use std::fmt;
|
use std::fmt;
|
||||||
|
|
||||||
/// Entity type classification (extensible enum).
|
/// Entity type classification (extensible enum).
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Hash)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Hash)]
|
||||||
#[serde(rename_all = "snake_case")]
|
#[serde(rename_all = "snake_case")]
|
||||||
pub enum EntityType {
|
pub enum EntityType {
|
||||||
Person,
|
Person,
|
||||||
@@ -17,6 +17,13 @@ pub enum EntityType {
|
|||||||
Location,
|
Location,
|
||||||
Event,
|
Event,
|
||||||
Organization,
|
Organization,
|
||||||
|
/// Agent prompt template tracked as a first-class entity.
|
||||||
|
/// Enables the agent to learn which prompts produce good results.
|
||||||
|
AgentPrompt,
|
||||||
|
/// Agent skill — a reusable capability the agent has learned.
|
||||||
|
AgentSkill,
|
||||||
|
/// Agent decision — a recorded choice with reasoning and outcome.
|
||||||
|
AgentDecision,
|
||||||
Unknown,
|
Unknown,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -29,6 +36,9 @@ impl EntityType {
|
|||||||
Self::Location => "location",
|
Self::Location => "location",
|
||||||
Self::Event => "event",
|
Self::Event => "event",
|
||||||
Self::Organization => "organization",
|
Self::Organization => "organization",
|
||||||
|
Self::AgentPrompt => "agent_prompt",
|
||||||
|
Self::AgentSkill => "agent_skill",
|
||||||
|
Self::AgentDecision => "agent_decision",
|
||||||
Self::Unknown => "unknown",
|
Self::Unknown => "unknown",
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -41,11 +51,24 @@ impl EntityType {
|
|||||||
"location" => Self::Location,
|
"location" => Self::Location,
|
||||||
"event" => Self::Event,
|
"event" => Self::Event,
|
||||||
"organization" => Self::Organization,
|
"organization" => Self::Organization,
|
||||||
|
"agent_prompt" => Self::AgentPrompt,
|
||||||
|
"agent_skill" => Self::AgentSkill,
|
||||||
|
"agent_decision" => Self::AgentDecision,
|
||||||
_ => Self::Unknown,
|
_ => Self::Unknown,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl<'de> serde::Deserialize<'de> for EntityType {
|
||||||
|
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
|
||||||
|
where
|
||||||
|
D: serde::Deserializer<'de>,
|
||||||
|
{
|
||||||
|
let s = String::deserialize(deserializer)?;
|
||||||
|
Ok(Self::from_str(&s))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
impl fmt::Display for EntityType {
|
impl fmt::Display for EntityType {
|
||||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||||
write!(f, "{}", self.as_str())
|
write!(f, "{}", self.as_str())
|
||||||
@@ -175,6 +198,9 @@ mod tests {
|
|||||||
EntityType::Person,
|
EntityType::Person,
|
||||||
EntityType::Tool,
|
EntityType::Tool,
|
||||||
EntityType::Concept,
|
EntityType::Concept,
|
||||||
|
EntityType::AgentPrompt,
|
||||||
|
EntityType::AgentSkill,
|
||||||
|
EntityType::AgentDecision,
|
||||||
] {
|
] {
|
||||||
let s = ty.as_str();
|
let s = ty.as_str();
|
||||||
assert_eq!(EntityType::from_str(s), *ty);
|
assert_eq!(EntityType::from_str(s), *ty);
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ pub mod scoring;
|
|||||||
pub mod entity;
|
pub mod entity;
|
||||||
pub mod edge;
|
pub mod edge;
|
||||||
pub mod community;
|
pub mod community;
|
||||||
|
pub mod agent_entity;
|
||||||
|
|
||||||
pub use gate_parser::{GateResponse, ParseError, parse_gate_response};
|
pub use gate_parser::{GateResponse, ParseError, parse_gate_response};
|
||||||
|
|
||||||
@@ -30,3 +31,4 @@ pub use scoring::{DocumentScorer, ScoringPipeline, GlobalTfIdfScorer, ProjectTfI
|
|||||||
pub use entity::{Entity, EntityType};
|
pub use entity::{Entity, EntityType};
|
||||||
pub use edge::{Edge, ContradictionStatus};
|
pub use edge::{Edge, ContradictionStatus};
|
||||||
pub use community::Community;
|
pub use community::Community;
|
||||||
|
pub use agent_entity::{AgentPromptMeta, AgentSkillMeta, AgentDecisionMeta, DecisionOutcome};
|
||||||
|
|||||||
@@ -52,13 +52,24 @@ impl AuthentikJwtIssuer {
|
|||||||
|
|
||||||
/// From environment: AUTHENTIK_ISSUER, AUTHENTIK_CLIENT_ID, AUTHENTIK_CLIENT_SECRET
|
/// From environment: AUTHENTIK_ISSUER, AUTHENTIK_CLIENT_ID, AUTHENTIK_CLIENT_SECRET
|
||||||
pub fn from_env() -> Result<Self> {
|
pub fn from_env() -> Result<Self> {
|
||||||
|
// Support both naming conventions: AUTHENTIK_* and memory-agent-oidc secret keys
|
||||||
let issuer = std::env::var("AUTHENTIK_ISSUER")
|
let issuer = std::env::var("AUTHENTIK_ISSUER")
|
||||||
.map_err(|_| anyhow!("AUTHENTIK_ISSUER not set"))?;
|
.or_else(|_| std::env::var("ISSUER"))
|
||||||
|
.map_err(|_| anyhow!("AUTHENTIK_ISSUER or ISSUER not set"))?;
|
||||||
let client_id = std::env::var("AUTHENTIK_CLIENT_ID")
|
let client_id = std::env::var("AUTHENTIK_CLIENT_ID")
|
||||||
.map_err(|_| anyhow!("AUTHENTIK_CLIENT_ID not set"))?;
|
.or_else(|_| std::env::var("CLIENT_ID"))
|
||||||
|
.map_err(|_| anyhow!("AUTHENTIK_CLIENT_ID or CLIENT_ID not set"))?;
|
||||||
let client_secret = std::env::var("AUTHENTIK_CLIENT_SECRET")
|
let client_secret = std::env::var("AUTHENTIK_CLIENT_SECRET")
|
||||||
.map_err(|_| anyhow!("AUTHENTIK_CLIENT_SECRET not set"))?;
|
.or_else(|_| std::env::var("CLIENT_SECRET"))
|
||||||
|
.map_err(|_| anyhow!("AUTHENTIK_CLIENT_SECRET or CLIENT_SECRET not set"))?;
|
||||||
|
|
||||||
|
tracing::info!(
|
||||||
|
target: "observability",
|
||||||
|
event = "authentik_jwt_init",
|
||||||
|
issuer = %issuer,
|
||||||
|
client_id = %client_id,
|
||||||
|
"Authentik JWT issuer initialized"
|
||||||
|
);
|
||||||
Ok(Self::new(&issuer, &client_id, &client_secret))
|
Ok(Self::new(&issuer, &client_id, &client_secret))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -92,12 +103,25 @@ impl AuthentikJwtIssuer {
|
|||||||
let client = reqwest::Client::new();
|
let client = reqwest::Client::new();
|
||||||
|
|
||||||
// Authentik OAuth2 token endpoint
|
// Authentik OAuth2 token endpoint
|
||||||
let token_url = format!("{}/token/", self.issuer_url.trim_end_matches('/'));
|
// Use TOKEN_URL env var if set, otherwise derive from issuer
|
||||||
|
let token_url = std::env::var("TOKEN_URL")
|
||||||
|
.or_else(|_| std::env::var("AUTHENTIK_TOKEN_URL"))
|
||||||
|
.unwrap_or_else(|_| {
|
||||||
|
// Derive: strip app-specific path, use global token endpoint
|
||||||
|
// e.g., https://authentik.riotpiao.com/application/o/memory-agent/
|
||||||
|
// -> https://authentik.riotpiao.com/application/o/token/
|
||||||
|
if let Some(base) = self.issuer_url.rfind("/o/") {
|
||||||
|
format!("{}/o/token/", &self.issuer_url[..base])
|
||||||
|
} else {
|
||||||
|
format!("{}/token/", self.issuer_url.trim_end_matches('/'))
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
let params = [
|
let params = [
|
||||||
("grant_type", "client_credentials"),
|
("grant_type", "client_credentials"),
|
||||||
("client_id", &self.client_id),
|
("client_id", &self.client_id),
|
||||||
("client_secret", &self.client_secret),
|
("client_secret", &self.client_secret),
|
||||||
|
("scope", "openid roles"),
|
||||||
];
|
];
|
||||||
|
|
||||||
let response = client
|
let response = client
|
||||||
|
|||||||
@@ -22,11 +22,15 @@ use tokio::sync::Mutex;
|
|||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||||
pub struct ExtractedEntity {
|
pub struct ExtractedEntity {
|
||||||
pub name: String,
|
pub name: String,
|
||||||
|
#[serde(alias = "type")]
|
||||||
pub entity_type: EntityType,
|
pub entity_type: EntityType,
|
||||||
pub summary: String,
|
pub summary: String,
|
||||||
|
#[serde(default = "default_confidence")]
|
||||||
pub confidence: f32,
|
pub confidence: f32,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn default_confidence() -> f32 { 0.8 }
|
||||||
|
|
||||||
impl ExtractedEntity {
|
impl ExtractedEntity {
|
||||||
/// Convert to domain model (Phase 1 type)
|
/// Convert to domain model (Phase 1 type)
|
||||||
pub fn to_domain(&self, project_id: &str) -> Entity {
|
pub fn to_domain(&self, project_id: &str) -> Entity {
|
||||||
@@ -62,6 +66,35 @@ impl LlmEntityExtractor {
|
|||||||
|
|
||||||
/// Parse extraction response JSON
|
/// Parse extraction response JSON
|
||||||
/// Format: { "entities": [{ "name": "...", "type": "...", "summary": "..." }, ...] }
|
/// Format: { "entities": [{ "name": "...", "type": "...", "summary": "..." }, ...] }
|
||||||
|
/// Clean LLM response: strip thinking tags, markdown fences, extract JSON
|
||||||
|
fn clean_llm_response(text: &str) -> String {
|
||||||
|
let mut result = text.to_string();
|
||||||
|
// Remove <think>...</think> blocks
|
||||||
|
while let Some(start) = result.find("<think>") {
|
||||||
|
if let Some(end) = result.find("</think>") {
|
||||||
|
result = format!("{}{}", &result[..start], &result[end + 8..]);
|
||||||
|
} else {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Remove markdown code fences
|
||||||
|
result = result.replace("```json", "").replace("```", "");
|
||||||
|
// Find JSON object
|
||||||
|
let trimmed = result.trim();
|
||||||
|
if let Some(start) = trimmed.find('{') {
|
||||||
|
if let Some(end) = trimmed.rfind('}') {
|
||||||
|
return trimmed[start..=end].to_string();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Maybe it's a JSON array — wrap in object
|
||||||
|
if let Some(start) = trimmed.find('[') {
|
||||||
|
if let Some(end) = trimmed.rfind(']') {
|
||||||
|
return format!("{{\"entities\": {}}}", &trimmed[start..=end]);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
trimmed.to_string()
|
||||||
|
}
|
||||||
|
|
||||||
fn parse_extraction(response: &str) -> Result<Vec<ExtractedEntity>> {
|
fn parse_extraction(response: &str) -> Result<Vec<ExtractedEntity>> {
|
||||||
#[derive(Deserialize)]
|
#[derive(Deserialize)]
|
||||||
struct Response {
|
struct Response {
|
||||||
@@ -123,7 +156,7 @@ impl LlmEntityExtractor {
|
|||||||
{"role": "user", "content": prompt}
|
{"role": "user", "content": prompt}
|
||||||
],
|
],
|
||||||
"temperature": 0.3,
|
"temperature": 0.3,
|
||||||
"max_tokens": 500
|
"max_tokens": 12000
|
||||||
});
|
});
|
||||||
|
|
||||||
let response = client
|
let response = client
|
||||||
@@ -131,7 +164,7 @@ impl LlmEntityExtractor {
|
|||||||
.header("Authorization", auth_header)
|
.header("Authorization", auth_header)
|
||||||
.header("Content-Type", "application/json")
|
.header("Content-Type", "application/json")
|
||||||
.json(&payload)
|
.json(&payload)
|
||||||
.timeout(std::time::Duration::from_secs(30))
|
.timeout(std::time::Duration::from_secs(90))
|
||||||
.send()
|
.send()
|
||||||
.await?;
|
.await?;
|
||||||
|
|
||||||
@@ -146,12 +179,28 @@ impl LlmEntityExtractor {
|
|||||||
}
|
}
|
||||||
|
|
||||||
let data: serde_json::Value = response.json().await?;
|
let data: serde_json::Value = response.json().await?;
|
||||||
let content = data["choices"][0]["message"]["content"]
|
// Extract content — some models put JSON in "content", others in "reasoning"
|
||||||
.as_str()
|
let msg = &data["choices"][0]["message"];
|
||||||
.unwrap_or("{}")
|
let raw_content = msg["content"].as_str().unwrap_or("").to_string();
|
||||||
.to_string();
|
let raw_reasoning = msg["reasoning"].as_str().unwrap_or("").to_string();
|
||||||
|
|
||||||
tracing::debug!("LLM response (via Authentik JWT): {}", content);
|
// Use content if non-empty, otherwise try reasoning field
|
||||||
|
let raw = if !raw_content.trim().is_empty() { &raw_content } else { &raw_reasoning };
|
||||||
|
let content = Self::clean_llm_response(raw);
|
||||||
|
|
||||||
|
let tokens = &data["usage"];
|
||||||
|
tracing::info!(
|
||||||
|
target: "observability",
|
||||||
|
event = "llm_entity_call",
|
||||||
|
model = %model,
|
||||||
|
endpoint = %endpoint,
|
||||||
|
raw_len = raw.len(),
|
||||||
|
cleaned_len = content.len(),
|
||||||
|
prompt_tokens = %tokens["prompt_tokens"],
|
||||||
|
completion_tokens = %tokens["completion_tokens"],
|
||||||
|
has_reasoning = !raw_reasoning.is_empty(),
|
||||||
|
"LLM entity extraction call complete"
|
||||||
|
);
|
||||||
Ok(content)
|
Ok(content)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -233,14 +282,27 @@ Respond in JSON:
|
|||||||
);
|
);
|
||||||
|
|
||||||
let reflection = if std::env::var("LLM_ENDPOINT").is_ok() {
|
let reflection = if std::env::var("LLM_ENDPOINT").is_ok() {
|
||||||
self.call_llm_endpoint(&reflection_prompt).await.unwrap_or_else(|_| self.simulate_llm(&reflection_prompt).unwrap_or_default())
|
self.call_llm_endpoint(&reflection_prompt).await.unwrap_or_else(|e| {
|
||||||
|
tracing::warn!("Reflection LLM call failed: {}, skipping verification", e);
|
||||||
|
String::new()
|
||||||
|
})
|
||||||
} else {
|
} else {
|
||||||
self.simulate_llm(&reflection_prompt)?
|
self.simulate_llm(&reflection_prompt)?
|
||||||
};
|
};
|
||||||
let verified = Self::parse_reflection(&reflection)?;
|
|
||||||
|
|
||||||
// Filter: keep only entities marked present
|
// If reflection succeeded, filter entities; otherwise keep all
|
||||||
entities.retain(|e| verified.iter().any(|(name, present)| name == &e.name && *present));
|
if !reflection.is_empty() {
|
||||||
|
match Self::parse_reflection(&reflection) {
|
||||||
|
Ok(verified) => {
|
||||||
|
entities.retain(|e| verified.iter().any(|(name, present)| name == &e.name && *present));
|
||||||
|
}
|
||||||
|
Err(e) => {
|
||||||
|
tracing::warn!("Reflection parse failed: {}, keeping all entities", e);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
tracing::info!("Reflection skipped, keeping {} unverified entities", entities.len());
|
||||||
|
}
|
||||||
|
|
||||||
// Adjust confidence for reflected entities (slight penalty for needing verification)
|
// Adjust confidence for reflected entities (slight penalty for needing verification)
|
||||||
for entity in &mut entities {
|
for entity in &mut entities {
|
||||||
|
|||||||
@@ -1,12 +1,12 @@
|
|||||||
//! Fact extraction: Identify relationships between entities
|
//! Fact extraction: Identify relationships between entities
|
||||||
//!
|
//!
|
||||||
//! Two implementations:
|
//! Three implementations:
|
||||||
//! 1. SimpleFactExtractor: Pattern-based (verbs + wiki links)
|
//! 1. SimpleFactExtractor: Pattern-based (verbs + wiki links)
|
||||||
//! 2. LlmFactExtractor: LLM-based (placeholder for production)
|
//! 2. LlmFactExtractor: LLM-based extraction with entity context
|
||||||
|
//! 3. Fallback chain: LLM → Simple pattern matching
|
||||||
//!
|
//!
|
||||||
//! CRAP: 12 (Simple pattern matching + LLM placeholder)
|
//! Aligned with Zep paper §2.2.2: Facts as edges between entity pairs,
|
||||||
//! SOLID: Trait-based (Open/Closed)
|
//! with temporal extraction and dedup against existing edges.
|
||||||
//! DRY: Reuses EntityExtractor pattern
|
|
||||||
|
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
@@ -27,20 +27,18 @@ pub struct ExtractedFact {
|
|||||||
pub trait FactExtractor: Send + Sync {
|
pub trait FactExtractor: Send + Sync {
|
||||||
async fn extract(&self, text: &str) -> Result<Vec<ExtractedFact>>;
|
async fn extract(&self, text: &str) -> Result<Vec<ExtractedFact>>;
|
||||||
|
|
||||||
/// Extract facts with GRM context (optional, defaults to extract())
|
/// Extract facts with entity context (Zep §2.2.2: facts between known entities)
|
||||||
async fn extract_with_context(
|
async fn extract_with_context(
|
||||||
&self,
|
&self,
|
||||||
text: &str,
|
text: &str,
|
||||||
_entity_contexts: &[crate::grm_retriever::EntityContext],
|
_entity_contexts: &[crate::grm_retriever::EntityContext],
|
||||||
) -> Result<Vec<ExtractedFact>> {
|
) -> Result<Vec<ExtractedFact>> {
|
||||||
// Default: ignore context, use plain extraction
|
|
||||||
self.extract(text).await
|
self.extract(text).await
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Simple fact extractor based on verb patterns
|
/// Simple fact extractor based on verb patterns
|
||||||
/// Pattern: [[Entity1]] verb [[Entity2]]
|
/// Pattern: [[Entity1]] verb [[Entity2]]
|
||||||
/// Common verbs: uses, manages, runs, deployed_to, works_with
|
|
||||||
pub struct SimpleFactExtractor;
|
pub struct SimpleFactExtractor;
|
||||||
|
|
||||||
#[async_trait]
|
#[async_trait]
|
||||||
@@ -48,17 +46,15 @@ impl FactExtractor for SimpleFactExtractor {
|
|||||||
async fn extract(&self, text: &str) -> Result<Vec<ExtractedFact>> {
|
async fn extract(&self, text: &str) -> Result<Vec<ExtractedFact>> {
|
||||||
let mut facts = vec![];
|
let mut facts = vec![];
|
||||||
|
|
||||||
// Extract [[Entity]] patterns
|
|
||||||
let entity_pattern = Regex::new(r"\[\[([^\]]+)\]\]")?;
|
let entity_pattern = Regex::new(r"\[\[([^\]]+)\]\]")?;
|
||||||
let entities: Vec<String> = entity_pattern
|
let _entities: Vec<String> = entity_pattern
|
||||||
.captures_iter(text)
|
.captures_iter(text)
|
||||||
.filter_map(|cap| cap.get(1).map(|m| m.as_str().to_string()))
|
.filter_map(|cap| cap.get(1).map(|m| m.as_str().to_string()))
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
// Common relationship verbs
|
let verbs = ["uses", "manages", "runs", "deployed_to", "works_with",
|
||||||
let verbs = ["uses", "manages", "runs", "deployed_to", "works_with"];
|
"depends_on", "contains", "extends", "implements", "connects_to"];
|
||||||
|
|
||||||
// Simple heuristic: if two entities appear close together with a verb between them
|
|
||||||
for verb in &verbs {
|
for verb in &verbs {
|
||||||
let pattern = format!(
|
let pattern = format!(
|
||||||
r"\[\[([^\]]+)\]\].*?{}.*?\[\[([^\]]+)\]\]",
|
r"\[\[([^\]]+)\]\].*?{}.*?\[\[([^\]]+)\]\]",
|
||||||
@@ -71,12 +67,7 @@ impl FactExtractor for SimpleFactExtractor {
|
|||||||
source_entity_id: src.as_str().to_string(),
|
source_entity_id: src.as_str().to_string(),
|
||||||
target_entity_id: tgt.as_str().to_string(),
|
target_entity_id: tgt.as_str().to_string(),
|
||||||
relation_type: verb.to_uppercase(),
|
relation_type: verb.to_uppercase(),
|
||||||
fact: format!(
|
fact: format!("{} {} {}", src.as_str(), verb, tgt.as_str()),
|
||||||
"{} {} {}",
|
|
||||||
src.as_str(),
|
|
||||||
verb,
|
|
||||||
tgt.as_str()
|
|
||||||
),
|
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -87,18 +78,251 @@ impl FactExtractor for SimpleFactExtractor {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// LLM-based fact extractor (placeholder for production)
|
/// LLM-based fact extractor (Zep §2.2.2 alignment)
|
||||||
/// TODO (Phase 2.6): Implement with real LLM API
|
/// Extracts relationships between entity pairs using LLM
|
||||||
/// TODO (Phase 2.6): Support complex relationships (3-way, temporal, conditional)
|
pub struct LlmFactExtractor {
|
||||||
pub struct LlmFactExtractor;
|
model_name: String,
|
||||||
|
jwt_issuer: Option<std::sync::Arc<tokio::sync::Mutex<crate::authentik_jwt::AuthentikJwtIssuer>>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl LlmFactExtractor {
|
||||||
|
pub fn new(model_name: &str) -> Self {
|
||||||
|
let jwt_issuer = crate::authentik_jwt::AuthentikJwtIssuer::from_env().ok();
|
||||||
|
Self {
|
||||||
|
model_name: model_name.to_string(),
|
||||||
|
jwt_issuer: jwt_issuer.map(|iss| std::sync::Arc::new(tokio::sync::Mutex::new(iss))),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Clean LLM response: strip thinking tags, markdown fences, extract JSON
|
||||||
|
fn clean_llm_response(text: &str) -> String {
|
||||||
|
let mut result = text.to_string();
|
||||||
|
while let Some(start) = result.find("<think>") {
|
||||||
|
if let Some(end) = result.find("</think>") {
|
||||||
|
result = format!("{}{}", &result[..start], &result[end + 8..]);
|
||||||
|
} else { break; }
|
||||||
|
}
|
||||||
|
result = result.replace("```json", "").replace("```", "");
|
||||||
|
let trimmed = result.trim();
|
||||||
|
if let Some(start) = trimmed.find('{') {
|
||||||
|
if let Some(end) = trimmed.rfind('}') {
|
||||||
|
return trimmed[start..=end].to_string();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if let Some(start) = trimmed.find('[') {
|
||||||
|
if let Some(end) = trimmed.rfind(']') {
|
||||||
|
return format!("{{\"facts\": {}}}", &trimmed[start..=end]);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
trimmed.to_string()
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn call_llm(&self, prompt: &str) -> Result<String> {
|
||||||
|
let endpoint = std::env::var("LLM_ENDPOINT")
|
||||||
|
.unwrap_or_else(|_| "http://localhost:11434/v1/chat/completions".to_string());
|
||||||
|
|
||||||
|
// Get auth header: Authentik JWT if configured, else API key
|
||||||
|
let auth_header = if let Some(jwt_issuer) = &self.jwt_issuer {
|
||||||
|
let issuer = jwt_issuer.lock().await;
|
||||||
|
match issuer.get_access_token().await {
|
||||||
|
Ok(token) => format!("Bearer {}", token),
|
||||||
|
Err(e) => {
|
||||||
|
tracing::warn!(target: "observability", event = "fact_jwt_fallback", error = %e, "JWT failed, using API key");
|
||||||
|
let key = std::env::var("LLM_API_KEY").unwrap_or_else(|_| "default-key".to_string());
|
||||||
|
format!("Bearer {}", key)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
let key = std::env::var("LLM_API_KEY")
|
||||||
|
.or_else(|_| std::env::var("MEM_API_KEY"))
|
||||||
|
.unwrap_or_else(|_| "default-key".to_string());
|
||||||
|
format!("Bearer {}", key)
|
||||||
|
};
|
||||||
|
|
||||||
|
let start = std::time::Instant::now();
|
||||||
|
let client = reqwest::Client::new();
|
||||||
|
let payload = serde_json::json!({
|
||||||
|
"model": self.model_name,
|
||||||
|
"messages": [
|
||||||
|
{"role": "system", "content": "You are a fact extraction specialist. Extract relationships between entities from text. Output ONLY valid JSON."},
|
||||||
|
{"role": "user", "content": prompt}
|
||||||
|
],
|
||||||
|
"max_tokens": 12000,
|
||||||
|
"temperature": 0.1
|
||||||
|
});
|
||||||
|
|
||||||
|
let response = client
|
||||||
|
.post(&endpoint)
|
||||||
|
.header("Authorization", &auth_header)
|
||||||
|
.header("Content-Type", "application/json")
|
||||||
|
.json(&payload)
|
||||||
|
.timeout(std::time::Duration::from_secs(120))
|
||||||
|
.send()
|
||||||
|
.await?;
|
||||||
|
|
||||||
|
let status = response.status();
|
||||||
|
if !status.is_success() {
|
||||||
|
let body = response.text().await.unwrap_or_default();
|
||||||
|
tracing::warn!(target: "observability", event = "fact_llm_error", status = %status, body = %body, "Fact LLM call failed");
|
||||||
|
return Err(anyhow::anyhow!("LLM API error: {}", status));
|
||||||
|
}
|
||||||
|
|
||||||
|
let elapsed = start.elapsed();
|
||||||
|
let data: serde_json::Value = response.json().await?;
|
||||||
|
|
||||||
|
// Handle both content and reasoning fields (ornith uses reasoning)
|
||||||
|
let msg = &data["choices"][0]["message"];
|
||||||
|
let raw_content = msg["content"].as_str().unwrap_or("").to_string();
|
||||||
|
let raw_reasoning = msg["reasoning"].as_str().unwrap_or("").to_string();
|
||||||
|
let raw = if !raw_content.trim().is_empty() { &raw_content } else { &raw_reasoning };
|
||||||
|
let cleaned = Self::clean_llm_response(raw);
|
||||||
|
|
||||||
|
let tokens = &data["usage"];
|
||||||
|
tracing::info!(
|
||||||
|
target: "observability",
|
||||||
|
event = "llm_fact_call",
|
||||||
|
model = %self.model_name,
|
||||||
|
endpoint = %endpoint,
|
||||||
|
raw_len = raw.len(),
|
||||||
|
cleaned_len = cleaned.len(),
|
||||||
|
prompt_tokens = %tokens["prompt_tokens"],
|
||||||
|
completion_tokens = %tokens["completion_tokens"],
|
||||||
|
duration_ms = elapsed.as_millis() as u64,
|
||||||
|
has_reasoning = !raw_reasoning.is_empty(),
|
||||||
|
"LLM fact extraction call complete"
|
||||||
|
);
|
||||||
|
Ok(cleaned)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[async_trait]
|
#[async_trait]
|
||||||
impl FactExtractor for LlmFactExtractor {
|
impl FactExtractor for LlmFactExtractor {
|
||||||
async fn extract(&self, _text: &str) -> Result<Vec<ExtractedFact>> {
|
async fn extract(&self, text: &str) -> Result<Vec<ExtractedFact>> {
|
||||||
// TODO (Phase 2.6): Implement LLM-based extraction
|
self.extract_with_context(text, &[]).await
|
||||||
// Pattern: Send text to api.riotpiao.com with prompt
|
}
|
||||||
// Parse response for [source, relation, target] tuples
|
|
||||||
Ok(vec![])
|
async fn extract_with_context(
|
||||||
|
&self,
|
||||||
|
text: &str,
|
||||||
|
entity_contexts: &[crate::grm_retriever::EntityContext],
|
||||||
|
) -> Result<Vec<ExtractedFact>> {
|
||||||
|
// Build entity list for prompt
|
||||||
|
let entity_names: Vec<&str> = entity_contexts
|
||||||
|
.iter()
|
||||||
|
.map(|e| e.entity_name.as_str())
|
||||||
|
.collect();
|
||||||
|
|
||||||
|
if entity_names.is_empty() {
|
||||||
|
tracing::debug!("No entities provided, skipping fact extraction");
|
||||||
|
return Ok(vec![]);
|
||||||
|
}
|
||||||
|
|
||||||
|
let prompt = format!(
|
||||||
|
r#"Extract relationships (facts) between these entities from the text.
|
||||||
|
|
||||||
|
Entities: {:?}
|
||||||
|
|
||||||
|
Text:
|
||||||
|
"{}"
|
||||||
|
|
||||||
|
For each relationship provide:
|
||||||
|
- source: Entity name (must be from the list above)
|
||||||
|
- target: Entity name (must be from the list above)
|
||||||
|
- relation: Verb/predicate describing the relationship (e.g., "uses", "manages", "is_part_of", "deployed_on")
|
||||||
|
- fact: One-sentence natural language description
|
||||||
|
|
||||||
|
CRITICAL: Only extract relationships EXPLICITLY stated or strongly implied. Source and target must both be from the entity list.
|
||||||
|
|
||||||
|
Respond in JSON:
|
||||||
|
{{"facts": [{{"source": "...", "target": "...", "relation": "...", "fact": "..."}}, ...]}}
|
||||||
|
"#,
|
||||||
|
entity_names, text
|
||||||
|
);
|
||||||
|
|
||||||
|
let llm_ok = std::env::var("LLM_ENDPOINT").is_ok();
|
||||||
|
let response = if llm_ok {
|
||||||
|
match self.call_llm(&prompt).await {
|
||||||
|
Ok(r) => r,
|
||||||
|
Err(e) => {
|
||||||
|
tracing::warn!("Fact extraction LLM failed: {}, returning empty", e);
|
||||||
|
return Ok(vec![]);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
tracing::debug!("LLM_ENDPOINT not set, skipping LLM fact extraction");
|
||||||
|
return Ok(vec![]);
|
||||||
|
};
|
||||||
|
|
||||||
|
// Parse response
|
||||||
|
#[derive(Deserialize)]
|
||||||
|
struct FactResponse {
|
||||||
|
facts: Vec<RawFact>,
|
||||||
|
}
|
||||||
|
#[derive(Deserialize)]
|
||||||
|
struct RawFact {
|
||||||
|
source: String,
|
||||||
|
target: String,
|
||||||
|
relation: String,
|
||||||
|
fact: String,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try parsing, if trailing chars error try trimming to valid JSON
|
||||||
|
let parsed = match serde_json::from_str::<FactResponse>(&response) {
|
||||||
|
Ok(r) => Ok(r),
|
||||||
|
Err(e) if e.to_string().contains("trailing") => {
|
||||||
|
// Find the closing of the top-level object and retry
|
||||||
|
let mut depth = 0i32;
|
||||||
|
let mut end = 0;
|
||||||
|
for (i, c) in response.char_indices() {
|
||||||
|
match c {
|
||||||
|
'{' | '[' => depth += 1,
|
||||||
|
'}' | ']' => { depth -= 1; if depth == 0 { end = i + 1; break; } },
|
||||||
|
_ => {}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if end > 0 {
|
||||||
|
serde_json::from_str::<FactResponse>(&response[..end])
|
||||||
|
} else {
|
||||||
|
Err(e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Err(e) => Err(e),
|
||||||
|
};
|
||||||
|
match parsed {
|
||||||
|
Ok(parsed) => {
|
||||||
|
let facts: Vec<ExtractedFact> = parsed.facts
|
||||||
|
.into_iter()
|
||||||
|
.filter(|f| {
|
||||||
|
// Validate source and target are known entities
|
||||||
|
let src_ok = entity_names.iter().any(|e| e.eq_ignore_ascii_case(&f.source));
|
||||||
|
let tgt_ok = entity_names.iter().any(|e| e.eq_ignore_ascii_case(&f.target));
|
||||||
|
if !src_ok || !tgt_ok {
|
||||||
|
tracing::debug!(
|
||||||
|
"Dropping fact with unknown entity: {} -> {}",
|
||||||
|
f.source, f.target
|
||||||
|
);
|
||||||
|
}
|
||||||
|
src_ok && tgt_ok && f.source != f.target
|
||||||
|
})
|
||||||
|
.map(|f| ExtractedFact {
|
||||||
|
source_entity_id: f.source,
|
||||||
|
target_entity_id: f.target,
|
||||||
|
relation_type: f.relation.to_uppercase(),
|
||||||
|
fact: f.fact,
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
|
||||||
|
tracing::info!(
|
||||||
|
"LLM fact extraction: {} facts from {} entities",
|
||||||
|
facts.len(), entity_names.len()
|
||||||
|
);
|
||||||
|
Ok(facts)
|
||||||
|
}
|
||||||
|
Err(e) => {
|
||||||
|
tracing::warn!("Fact extraction JSON parse failed: {}", e);
|
||||||
|
Ok(vec![])
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -110,9 +334,38 @@ mod tests {
|
|||||||
async fn test_simple_fact_extraction() {
|
async fn test_simple_fact_extraction() {
|
||||||
let extractor = SimpleFactExtractor;
|
let extractor = SimpleFactExtractor;
|
||||||
let text = "[[Rock]] uses [[Kubernetes]] and [[ArgoCD]]";
|
let text = "[[Rock]] uses [[Kubernetes]] and [[ArgoCD]]";
|
||||||
|
|
||||||
let facts = extractor.extract(text).await.unwrap();
|
let facts = extractor.extract(text).await.unwrap();
|
||||||
assert!(facts.len() > 0);
|
assert!(!facts.is_empty());
|
||||||
assert!(facts.iter().any(|f| f.relation_type == "USES"));
|
assert!(facts.iter().any(|f| f.relation_type == "USES"));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_simple_no_wiki_links() {
|
||||||
|
let extractor = SimpleFactExtractor;
|
||||||
|
let text = "Kubernetes uses etcd for storage";
|
||||||
|
let facts = extractor.extract(text).await.unwrap();
|
||||||
|
assert!(facts.is_empty()); // No [[wiki links]]
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_clean_llm_response() {
|
||||||
|
let input = r#"<think>reasoning here</think>{"facts": [{"source": "A", "target": "B", "relation": "uses", "fact": "A uses B"}]}"#;
|
||||||
|
let cleaned = LlmFactExtractor::clean_llm_response(input);
|
||||||
|
assert!(cleaned.starts_with("{"));
|
||||||
|
assert!(cleaned.contains("facts"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_strip_thinking_no_tags() {
|
||||||
|
let input = r#"{"facts": []}"#;
|
||||||
|
let cleaned = LlmFactExtractor::clean_llm_response(input);
|
||||||
|
assert_eq!(cleaned, input);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_llm_fact_no_entities_returns_empty() {
|
||||||
|
let extractor = LlmFactExtractor::new("test");
|
||||||
|
let facts = extractor.extract_with_context("some text", &[]).await.unwrap();
|
||||||
|
assert!(facts.is_empty());
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,67 @@
|
|||||||
|
-- Migration 009: Temporal edge schema (Zep paper §2.2.2)
|
||||||
|
-- Replaces old memory_edge (child_sha/parent_sha node graph)
|
||||||
|
-- with temporal edge schema supporting relation types, facts, and validity periods.
|
||||||
|
-- Idempotent: safe to run multiple times.
|
||||||
|
|
||||||
|
-- Rename old table if it still exists (skip if already migrated)
|
||||||
|
DO $$
|
||||||
|
BEGIN
|
||||||
|
IF EXISTS (SELECT 1 FROM information_schema.tables WHERE table_name = 'memory_edge'
|
||||||
|
AND EXISTS (SELECT 1 FROM information_schema.columns
|
||||||
|
WHERE table_name = 'memory_edge' AND column_name = 'child_sha'))
|
||||||
|
THEN
|
||||||
|
ALTER TABLE memory_edge RENAME TO memory_edge_legacy;
|
||||||
|
END IF;
|
||||||
|
END $$;
|
||||||
|
|
||||||
|
-- Create temporal edge table
|
||||||
|
CREATE TABLE IF NOT EXISTS memory_edge (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
project_id TEXT NOT NULL DEFAULT 'default',
|
||||||
|
source_id TEXT NOT NULL,
|
||||||
|
target_id TEXT NOT NULL,
|
||||||
|
relation_type TEXT NOT NULL DEFAULT '',
|
||||||
|
fact TEXT NOT NULL DEFAULT '',
|
||||||
|
weight REAL NOT NULL DEFAULT 1.0,
|
||||||
|
strength REAL DEFAULT 1.0,
|
||||||
|
confidence REAL DEFAULT 0.8,
|
||||||
|
t_valid TIMESTAMPTZ,
|
||||||
|
t_invalid TIMESTAMPTZ,
|
||||||
|
t_created TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||||
|
t_expired TIMESTAMPTZ,
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||||
|
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||||
|
episode_id TEXT,
|
||||||
|
deleted_at TIMESTAMPTZ
|
||||||
|
);
|
||||||
|
|
||||||
|
-- Ensure app user owns the table
|
||||||
|
DO $$ BEGIN
|
||||||
|
IF EXISTS (SELECT 1 FROM pg_roles WHERE rolname = 'app') THEN
|
||||||
|
ALTER TABLE memory_edge OWNER TO app;
|
||||||
|
END IF;
|
||||||
|
END $$;
|
||||||
|
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_memory_edge_source ON memory_edge(source_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_memory_edge_target ON memory_edge(target_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_memory_edge_project ON memory_edge(project_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_memory_edge_relation ON memory_edge(relation_type);
|
||||||
|
|
||||||
|
-- Ensure memory_entity has all columns code expects
|
||||||
|
ALTER TABLE memory_entity ADD COLUMN IF NOT EXISTS deleted_at TIMESTAMPTZ;
|
||||||
|
ALTER TABLE memory_entity ADD COLUMN IF NOT EXISTS source_count INTEGER DEFAULT 1;
|
||||||
|
|
||||||
|
-- Unique constraint for entity upsert dedup
|
||||||
|
DO $$
|
||||||
|
BEGIN
|
||||||
|
-- Dedup existing rows before creating unique index
|
||||||
|
DELETE FROM memory_entity a USING memory_entity b
|
||||||
|
WHERE a.project_id = b.project_id AND a.name = b.name
|
||||||
|
AND a.t_created < b.t_created;
|
||||||
|
EXCEPTION WHEN OTHERS THEN NULL;
|
||||||
|
END $$;
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS idx_memory_entity_project_name ON memory_entity(project_id, name);
|
||||||
|
|
||||||
|
-- ROLLBACK instructions:
|
||||||
|
-- DROP TABLE IF EXISTS memory_edge;
|
||||||
|
-- ALTER TABLE IF EXISTS memory_edge_legacy RENAME TO memory_edge;
|
||||||
@@ -20,7 +20,6 @@ data:
|
|||||||
# OpenSearch
|
# OpenSearch
|
||||||
OPENSEARCH_HOST: "opensearch.poimen.svc.cluster.local:9200"
|
OPENSEARCH_HOST: "opensearch.poimen.svc.cluster.local:9200"
|
||||||
# Obsidian
|
# Obsidian
|
||||||
OBSIDIAN_URL: "http://obsidian-server.poimen.svc.cluster.local:8080"
|
|
||||||
# LLM Configuration (for entity extraction)
|
# LLM Configuration (for entity extraction)
|
||||||
LLM_ENDPOINT: "http://api-internal.riotpiao.com:8000/v1/chat/completions"
|
LLM_ENDPOINT: "http://api-internal.riotpiao.com:8000/v1/chat/completions"
|
||||||
LLM_MODEL: "qwen:7b"
|
LLM_MODEL: "qwen:7b"
|
||||||
|
|||||||
+35
-9
@@ -1,6 +1,6 @@
|
|||||||
# Poimen Memory API Server
|
# Poimen Memory API Server
|
||||||
# Serves 7 HTTP endpoints for memory ingest, query, and management.
|
# Serves HTTP endpoints for memory ingest, query, visualization.
|
||||||
# Connects to memory-db (pgvector) for persistent storage.
|
# Connects to memory-db (pgvector) + api.riotpiao.com (LLM via Authentik JWT).
|
||||||
apiVersion: apps/v1
|
apiVersion: apps/v1
|
||||||
kind: Deployment
|
kind: Deployment
|
||||||
metadata:
|
metadata:
|
||||||
@@ -60,13 +60,43 @@ spec:
|
|||||||
key: password
|
key: password
|
||||||
- name: DATABASE_URL
|
- name: DATABASE_URL
|
||||||
value: "postgresql://$(DATABASE_USER):$(DATABASE_PASSWORD)@$(DATABASE_HOST):$(DATABASE_PORT)/$(DATABASE_NAME)?sslmode=disable"
|
value: "postgresql://$(DATABASE_USER):$(DATABASE_PASSWORD)@$(DATABASE_HOST):$(DATABASE_PORT)/$(DATABASE_NAME)?sslmode=disable"
|
||||||
# LLM Gateway API key
|
|
||||||
|
# LLM via api.riotpiao.com (Authentik JWT auth)
|
||||||
|
- name: LLM_ENDPOINT
|
||||||
|
value: "https://api.riotpiao.com/v1/chat/completions"
|
||||||
|
- name: LLM_API_BASE
|
||||||
|
value: "https://api.riotpiao.com/v1"
|
||||||
|
- name: LLM_MODEL
|
||||||
|
value: "ornith:35b"
|
||||||
|
|
||||||
|
# Authentik service account (memory-agent-oidc secret)
|
||||||
|
- name: AUTHENTIK_ISSUER
|
||||||
|
valueFrom:
|
||||||
|
secretKeyRef:
|
||||||
|
name: memory-agent-oidc
|
||||||
|
key: ISSUER
|
||||||
|
- name: AUTHENTIK_CLIENT_ID
|
||||||
|
valueFrom:
|
||||||
|
secretKeyRef:
|
||||||
|
name: memory-agent-oidc
|
||||||
|
key: CLIENT_ID
|
||||||
|
- name: AUTHENTIK_CLIENT_SECRET
|
||||||
|
valueFrom:
|
||||||
|
secretKeyRef:
|
||||||
|
name: memory-agent-oidc
|
||||||
|
key: CLIENT_SECRET
|
||||||
|
- name: TOKEN_URL
|
||||||
|
valueFrom:
|
||||||
|
secretKeyRef:
|
||||||
|
name: memory-agent-oidc
|
||||||
|
key: TOKEN_URL
|
||||||
|
|
||||||
|
# Server config
|
||||||
- name: MEM_API_KEY
|
- name: MEM_API_KEY
|
||||||
valueFrom:
|
valueFrom:
|
||||||
secretKeyRef:
|
secretKeyRef:
|
||||||
name: poimen-memory-secrets
|
name: poimen-memory-secrets
|
||||||
key: llm-api-key
|
key: llm-api-key
|
||||||
# Server config (from ConfigMap)
|
|
||||||
- name: MEM_PORT
|
- name: MEM_PORT
|
||||||
value: "8080"
|
value: "8080"
|
||||||
- name: MEM_HOME
|
- name: MEM_HOME
|
||||||
@@ -74,10 +104,7 @@ spec:
|
|||||||
envFrom:
|
envFrom:
|
||||||
- configMapRef:
|
- configMapRef:
|
||||||
name: poimen-memory-config
|
name: poimen-memory-config
|
||||||
- secretRef:
|
command: ["/app/mem"]
|
||||||
name: poimen-memory-auth
|
|
||||||
- secretRef:
|
|
||||||
name: poimen-memory-secrets
|
|
||||||
args:
|
args:
|
||||||
- serve
|
- serve
|
||||||
- --port
|
- --port
|
||||||
@@ -110,7 +137,6 @@ spec:
|
|||||||
- name: tmp
|
- name: tmp
|
||||||
emptyDir:
|
emptyDir:
|
||||||
sizeLimit: 64Mi
|
sizeLimit: 64Mi
|
||||||
# Tolerate control-plane nodes
|
|
||||||
tolerations:
|
tolerations:
|
||||||
- key: node-role.kubernetes.io/control-plane
|
- key: node-role.kubernetes.io/control-plane
|
||||||
operator: Exists
|
operator: Exists
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ resources:
|
|||||||
- deployment.yaml
|
- deployment.yaml
|
||||||
- service.yaml
|
- service.yaml
|
||||||
- config.yaml
|
- config.yaml
|
||||||
- obsidian.yaml
|
# obsidian.yaml retired — reference docs now via memory graph
|
||||||
# Legacy secret managed separately
|
# Legacy secret managed separately
|
||||||
# - secrets.yaml
|
# - secrets.yaml
|
||||||
generators:
|
generators:
|
||||||
|
|||||||
@@ -1,159 +0,0 @@
|
|||||||
---
|
|
||||||
# Obsidian server deployment
|
|
||||||
# Serves local vault with web UI and API
|
|
||||||
apiVersion: apps/v1
|
|
||||||
kind: Deployment
|
|
||||||
metadata:
|
|
||||||
name: obsidian-server
|
|
||||||
namespace: poimen
|
|
||||||
labels:
|
|
||||||
app.kubernetes.io/name: obsidian-server
|
|
||||||
app.kubernetes.io/part-of: poimen-memory
|
|
||||||
spec:
|
|
||||||
replicas: 1
|
|
||||||
selector:
|
|
||||||
matchLabels:
|
|
||||||
app.kubernetes.io/name: obsidian-server
|
|
||||||
template:
|
|
||||||
metadata:
|
|
||||||
labels:
|
|
||||||
app.kubernetes.io/name: obsidian-server
|
|
||||||
app.kubernetes.io/part-of: poimen-memory
|
|
||||||
spec:
|
|
||||||
serviceAccountName: obsidian-server
|
|
||||||
securityContext:
|
|
||||||
runAsNonRoot: true
|
|
||||||
runAsUser: 1000
|
|
||||||
runAsGroup: 1000
|
|
||||||
fsGroup: 1000
|
|
||||||
seccompProfile:
|
|
||||||
type: RuntimeDefault
|
|
||||||
initContainers:
|
|
||||||
- name: git-sync-init
|
|
||||||
image: alpine/git:latest
|
|
||||||
securityContext:
|
|
||||||
runAsNonRoot: false
|
|
||||||
runAsUser: 0
|
|
||||||
allowPrivilegeEscalation: false
|
|
||||||
capabilities:
|
|
||||||
drop:
|
|
||||||
- ALL
|
|
||||||
add:
|
|
||||||
- CHOWN
|
|
||||||
- DAC_OVERRIDE
|
|
||||||
command:
|
|
||||||
- sh
|
|
||||||
- -c
|
|
||||||
- |
|
|
||||||
export GIT_SSH_COMMAND="ssh -i /root/.ssh/id_ed25519 -o StrictHostKeyChecking=no"
|
|
||||||
git config --global --add safe.directory /vault
|
|
||||||
if [ -d /vault/.git ]; then
|
|
||||||
cd /vault && git pull origin main || true
|
|
||||||
else
|
|
||||||
# Clone into temp, move contents into vault
|
|
||||||
rm -rf /tmp/repo
|
|
||||||
git clone ssh://[email protected]:2222/rock/poimen-obesdient-memory.git /tmp/repo
|
|
||||||
cp -a /tmp/repo/. /vault/
|
|
||||||
rm -rf /tmp/repo
|
|
||||||
fi
|
|
||||||
chown -R 1000:1000 /vault
|
|
||||||
volumeMounts:
|
|
||||||
- name: vault
|
|
||||||
mountPath: /vault
|
|
||||||
- name: ssh-key
|
|
||||||
mountPath: /root/.ssh
|
|
||||||
readOnly: true
|
|
||||||
containers:
|
|
||||||
- name: obsidian-server
|
|
||||||
image: ppatlabs/obsidian:latest
|
|
||||||
imagePullPolicy: IfNotPresent
|
|
||||||
securityContext:
|
|
||||||
allowPrivilegeEscalation: false
|
|
||||||
capabilities:
|
|
||||||
drop:
|
|
||||||
- ALL
|
|
||||||
ports:
|
|
||||||
- name: http
|
|
||||||
containerPort: 27124
|
|
||||||
protocol: TCP
|
|
||||||
env:
|
|
||||||
- name: VAULT_NAME
|
|
||||||
value: poimen-vault
|
|
||||||
- name: VAULT_PATH
|
|
||||||
value: /vault
|
|
||||||
- name: REST_API_ENABLED
|
|
||||||
value: "true"
|
|
||||||
- name: REST_API_PORT
|
|
||||||
value: "8080"
|
|
||||||
volumeMounts:
|
|
||||||
- name: vault
|
|
||||||
mountPath: /vault
|
|
||||||
- name: config
|
|
||||||
mountPath: /config
|
|
||||||
resources:
|
|
||||||
requests:
|
|
||||||
cpu: 100m
|
|
||||||
memory: 256Mi
|
|
||||||
limits:
|
|
||||||
cpu: 500m
|
|
||||||
memory: 512Mi
|
|
||||||
livenessProbe:
|
|
||||||
httpGet:
|
|
||||||
path: /
|
|
||||||
port: http
|
|
||||||
scheme: HTTPS
|
|
||||||
initialDelaySeconds: 30
|
|
||||||
periodSeconds: 10
|
|
||||||
timeoutSeconds: 5
|
|
||||||
readinessProbe:
|
|
||||||
httpGet:
|
|
||||||
path: /
|
|
||||||
port: http
|
|
||||||
scheme: HTTPS
|
|
||||||
initialDelaySeconds: 15
|
|
||||||
periodSeconds: 5
|
|
||||||
timeoutSeconds: 5
|
|
||||||
volumes:
|
|
||||||
- name: vault
|
|
||||||
persistentVolumeClaim:
|
|
||||||
claimName: obsidian-vault
|
|
||||||
- name: config
|
|
||||||
emptyDir: {}
|
|
||||||
- name: ssh-key
|
|
||||||
secret:
|
|
||||||
secretName: obsidian-git-ssh
|
|
||||||
defaultMode: 0400
|
|
||||||
|
|
||||||
# PVC managed by homelab repo (k8s/infra/databases/obsidian-vault-pvc.yaml)
|
|
||||||
|
|
||||||
---
|
|
||||||
# Service for Obsidian server
|
|
||||||
apiVersion: v1
|
|
||||||
kind: Service
|
|
||||||
metadata:
|
|
||||||
name: obsidian-server
|
|
||||||
namespace: poimen
|
|
||||||
labels:
|
|
||||||
app.kubernetes.io/name: obsidian-server
|
|
||||||
spec:
|
|
||||||
type: ClusterIP
|
|
||||||
ports:
|
|
||||||
- name: http
|
|
||||||
port: 80
|
|
||||||
targetPort: 27124
|
|
||||||
protocol: TCP
|
|
||||||
selector:
|
|
||||||
app.kubernetes.io/name: obsidian-server
|
|
||||||
|
|
||||||
---
|
|
||||||
# ServiceAccount for Obsidian
|
|
||||||
apiVersion: v1
|
|
||||||
kind: ServiceAccount
|
|
||||||
metadata:
|
|
||||||
name: obsidian-server
|
|
||||||
namespace: poimen
|
|
||||||
labels:
|
|
||||||
app.kubernetes.io/name: obsidian-server
|
|
||||||
|
|
||||||
# Ingress managed by homelab repo (obsidian.riotpiao.com)
|
|
||||||
# See: homelab/k8s/bootstrap/ingress/ingress.yaml
|
|
||||||
@@ -1,8 +1,6 @@
|
|||||||
---
|
# Dedicated CNPG Postgres for Poimen Memory (GitOps, wave 2).
|
||||||
# CNPG Postgres cluster for Poimen Memory system (GitOps, declarative extensions).
|
# Matches homelab/k8s/infra/databases/memory-db.yaml — single source of truth.
|
||||||
# 2 instances, pgvector 0.7.0 via spec.extensions (not manual CREATE EXTENSION).
|
# CNPG generates secret `memory-db-app` + service `memory-db-rw` in ns poimen.
|
||||||
# Storage: 10Gi longhorn, consistent with temporal-db.yaml.
|
|
||||||
# No manual psql needed — all via git/ArgoCD.
|
|
||||||
apiVersion: postgresql.cnpg.io/v1
|
apiVersion: postgresql.cnpg.io/v1
|
||||||
kind: Cluster
|
kind: Cluster
|
||||||
metadata:
|
metadata:
|
||||||
@@ -13,20 +11,6 @@ metadata:
|
|||||||
spec:
|
spec:
|
||||||
instances: 2
|
instances: 2
|
||||||
imageName: ghcr.io/cloudnative-pg/postgresql:16.2
|
imageName: ghcr.io/cloudnative-pg/postgresql:16.2
|
||||||
enableSuperuserAccess: false
|
|
||||||
storage:
|
|
||||||
size: 10Gi
|
|
||||||
storageClass: longhorn
|
|
||||||
resources:
|
|
||||||
requests: { memory: "512Mi", cpu: "250m" }
|
|
||||||
limits: { memory: "2Gi", cpu: "1" }
|
|
||||||
affinity:
|
|
||||||
podAntiAffinityType: preferred
|
|
||||||
topologyKey: kubernetes.io/hostname
|
|
||||||
tolerations:
|
|
||||||
- key: node-role.kubernetes.io/control-plane
|
|
||||||
operator: Exists
|
|
||||||
effect: NoSchedule
|
|
||||||
bootstrap:
|
bootstrap:
|
||||||
initdb:
|
initdb:
|
||||||
database: memory
|
database: memory
|
||||||
@@ -34,25 +18,19 @@ spec:
|
|||||||
encoding: UTF8
|
encoding: UTF8
|
||||||
localeCollate: C
|
localeCollate: C
|
||||||
localeCType: C
|
localeCType: C
|
||||||
monitoring:
|
postInitApplicationSQL:
|
||||||
enabled: true
|
- "CREATE EXTENSION vector;"
|
||||||
podMonitorTemplate:
|
enableSuperuserAccess: false
|
||||||
spec:
|
resources:
|
||||||
interval: 30s
|
requests: { memory: "512Mi", cpu: "250m" }
|
||||||
scrapeTimeout: 10s
|
limits: { memory: "2Gi", cpu: "1" }
|
||||||
---
|
storage:
|
||||||
# Database resource with pgvector extension (declarative, git-managed).
|
size: 20Gi
|
||||||
# CNPG 1.30.0+ supports this via spec.extensions on the Database CRD.
|
storageClass: longhorn
|
||||||
# Ensures pgvector is installed and available for HNSW indexing.
|
affinity:
|
||||||
apiVersion: postgresql.cnpg.io/v1
|
podAntiAffinityType: preferred
|
||||||
kind: Database
|
topologyKey: kubernetes.io/hostname
|
||||||
metadata:
|
tolerations:
|
||||||
name: memory
|
- key: node-role.kubernetes.io/control-plane
|
||||||
namespace: poimen
|
operator: Exists
|
||||||
spec:
|
effect: NoSchedule
|
||||||
cluster:
|
|
||||||
name: memory-db
|
|
||||||
owner: app
|
|
||||||
extensions:
|
|
||||||
- name: vector
|
|
||||||
ensure: present
|
|
||||||
|
|||||||
Reference in New Issue
Block a user