diff --git a/.github/workflows/code-review.yml b/.github/workflows/code-review.yml new file mode 100644 index 0000000..de812b9 --- /dev/null +++ b/.github/workflows/code-review.yml @@ -0,0 +1,87 @@ +name: AI Code Review + +on: + pull_request: + types: [opened, synchronize] + +permissions: + contents: read + pull-requests: write + +jobs: + review: + runs-on: ubuntu-latest + timeout-minutes: 15 + steps: + - uses: actions/checkout@v4 + with: + fetch-depth: 0 # Full history for diff + + - uses: actions/setup-node@v4 + with: + node-version: '20' + + - name: Install Open Code Review + run: npm install -g @alibaba-group/open-code-review + + - name: Configure OCR with MiMo + env: + MIMO_API_KEY: ${{ secrets.MIMO_API_KEY }} + MIMO_BASE_URL: ${{ secrets.MIMO_BASE_URL }} + run: | + ocr config set llm.url "$MIMO_BASE_URL" + ocr config set llm.auth_token "$MIMO_API_KEY" + ocr config set llm.model "mimo-v2.5" + ocr config set llm.use_anthropic false + ocr config set language "English" + + - name: Run Code Review + id: review + env: + GH_TOKEN: ${{ github.token }} + PR_TITLE: ${{ github.event.pull_request.title }} + PR_NUMBER: ${{ github.event.pull_request.number }} + BASE_SHA: ${{ github.event.pull_request.base.sha }} + HEAD_SHA: ${{ github.event.pull_request.head.sha }} + run: | + echo "Reviewing: $BASE_SHA..$HEAD_SHA" + + # Run OCR and capture output + REVIEW=$(ocr review \ + --from "$BASE_SHA" \ + --to "$HEAD_SHA" \ + --audience agent \ + --background "PR #${PR_NUMBER}: ${PR_TITLE}" \ + 2>&1) || true + + # Save to file for the comment step + echo "$REVIEW" > /tmp/ocr-review.txt + + # Check if there are actual comments + if echo "$REVIEW" | grep -q "comment(s)"; then + echo "has_issues=true" >> "$GITHUB_OUTPUT" + else + echo "has_issues=false" >> "$GITHUB_OUTPUT" + fi + + - name: Post Review Comment + if: always() + env: + GH_TOKEN: ${{ github.token }} + PR_NUMBER: ${{ github.event.pull_request.number }} + run: | + REVIEW=$(cat /tmp/ocr-review.txt) + + gh pr comment "$PR_NUMBER" \ + --body "## 🤖 AI Code Review (Open Code Review + MiMo) + +
+ Review Results + + \`\`\` + ${REVIEW} + \`\`\` + +
+ + _Automated review by [Open Code Review](https://github.com/alibaba/open-code-review) + Xiaomi MiMo-V2.5_" diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md new file mode 100644 index 0000000..e32f551 --- /dev/null +++ b/CONTRIBUTING.md @@ -0,0 +1,70 @@ +# Contributing to TicketPilot + +Thanks for your interest in contributing! Here's how to get started. + +## Development Setup + +```bash +git clone https://github.com/lennney/ticketpilot.git +cd ticketpilot +pip install uv +uv sync --group dev +docker compose up -d db +``` + +## Running Tests + +```bash +# Unit tests (no DB required, fast) +TICKETPILOT_SKIP_DB_TESTS=1 uv run pytest tests/ --ignore=tests/integration -q + +# Full tests (requires DB) +uv run pytest tests/ -v + +# Quality gate (must pass before PR) +bash scripts/run_quality_gate.sh +``` + +## Code Style + +- **Linter**: ruff (all rules enabled, no isort) +- **Type hints**: Required for all public functions +- **Docstrings**: Required for all public modules and classes +- **Tests**: Every new feature needs tests; coverage must stay ≥ 70% + +## How to Contribute + +### Reporting Bugs + +Open an issue with: +- Steps to reproduce +- Expected vs actual behavior +- Python version and OS + +### Submitting Changes + +1. Fork the repo +2. Create a branch: `git checkout -b feature/your-feature` +3. Make your changes with tests +4. Run the quality gate: `bash scripts/run_quality_gate.sh` +5. Submit a PR with a clear description + +### Good First Issues + +Look for issues labeled `good-first-issue`: + +- 📝 Documentation improvements +- 🧪 Test coverage for edge cases +- 🔧 Small bug fixes +- 🌐 Internationalization + +## Architecture Notes + +The pipeline is **deterministic by design** — no LLM calls in the core pipeline +(classification, risk, retrieval, confidence scoring). LLM is only used in +`DraftAgent` for reply generation. This is intentional: it means the pipeline +is fully testable without mocking LLM responses. + +## Questions? + +Open a discussion or comment on an existing issue. diff --git a/README.md b/README.md index 2afc257..61a2aa3 100644 --- a/README.md +++ b/README.md @@ -1,211 +1,182 @@ -# TicketPilot +# 🎫 TicketPilot -AI Customer Service Copilot for cross-border e-commerce — **deterministic, no-LLM-in-pipeline, full-chain traceability**. +**中文客服工单 AI 分拣系统 — 确定性管线,零 LLM 调用,全链路可追溯** -> TicketPilot chains intent classification, risk assessment, evidence retrieval, draft generation, and human review into a single pipeline. Human agents only need to judge the ~20% of tickets that actually require judgment. +> 跨境电商客服 Copilot:意图分类 → 风险评估 → 混合检索 → 证据化草稿 → 人工审核台 +> 60% 工单自动发送,40% 路由到人工,0% 关键工单遗漏 -## What Makes TicketPilot Different +[![Tests](https://img.shields.io/badge/tests-1%2C760-brightgreen)]() +[![Coverage](https://img.shields.io/badge/coverage-87%25-brightgreen)]() +[![Python](https://img.shields.io/badge/python-3.11%2B-blue)]() +[![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](LICENSE) +[![Docker](https://img.shields.io/badge/docker-compose-blue?logo=docker)]() -| Feature | Typical Approach | TicketPilot | -|---------|-----------------|-------------| -| Confidence scoring | Binary (confident / not) | 4-dimensional weighted: retrieval + classification + citation + evidence density | -| Response routing | All-auto or all-human | 4-tier degradation: AUTO_SEND → CAUTIOUS → HUMAN_REVIEW → ESCALATION | -| Hallucination guard | None or simple keyword filter | 8-category forbidden promise detection (refund amounts, legal threats, etc.) | -| Retrieval | Simple vector search | Keyword FTS + Vector HNSW → RRF fusion with per-ranker contribution tracing | -| Traceability | None | Full chain: answer → citation → chunk → document (ClaimProvenance + RetrievalTrace) | -| Agent architecture | Single agent | Multi-agent orchestrator with intent-based routing to 5 specialized agents | -| Pipeline determinism | LLM-dependent | Rule-driven, zero LLM calls in pipeline | -| Calibration | Static thresholds | Feedback loop with isotonic regression calibration + reliability diagrams | -| Experimentation | Manual A/B | Built-in A/B experiment framework with comparison reports | -| Self-reflection | None | Skill seed learning from successful draft patterns | +--- -## Architecture +## 为什么做这个项目 -```mermaid -graph TD - A[Ticket Input] --> B[Intent Classifier] - B --> C[Risk Assessor] - C --> D[Hybrid Retrieval
FTS + pgvector → RRF] - D --> E[Multi-Agent Router] - E --> F[Refund Agent] - E --> G[Complaint Agent] - E --> H[Logistics Agent] - E --> I[Technical Agent] - E --> J[Default Agent] - F --> K[Draft Generator] - G --> K - H --> K - I --> K - J --> K - K --> L[Claim Guard] - K --> M[Citation Validator] - L --> N[Confidence Scorer
4 dimensions] - M --> N - N --> O{Confidence Tier} - O -->|HIGH| P[Auto-Send] - O -->|MEDIUM| P - O -->|LOW| Q[Human Review] - O -->|CRITICAL| R[Force Escalation] - Q --> S{Decision} - S -->|Approve| P - S -->|Edit| T[Revise] - T --> P - P --> U[Feedback Loop] - U --> V[Isotonic Calibrator] - V --> B -``` - -## Key Modules - -### Confidence & Routing -- **ConfidenceScorer** — 4-dimensional scoring (retrieval 35%, classification 25%, citation 25%, evidence density 15%) -- **DegradationRouter** — 4-tier routing based on confidence level -- **Claim Guard** — Forbidden promise detection, citation coverage, risk acknowledgment -- **Citation Validator** — Luhn bank card check, unsupported claim detection - -### Multi-Agent System -- **Orchestrator** — Intent-based routing to specialized agents -- **5 Specialists** — RefundAgent, ComplaintAgent, LogisticsAgent, TechnicalAgent, DefaultAgent -- **Per-agent prompt templates** — Each specialist uses domain-specific prompts -- **Self-Reflection Skills** — Agents learn from successful draft patterns via skill seed - -### Retrieval -- **Hybrid search** — PostgreSQL FTS + pgvector HNSW → RRF fusion -- **RetrievalTrace** — Full explainability: keyword rank, vector rank, RRF contribution per result -- **Context truncation** — Token-budget-aware truncation for retrieval results - -### Feedback & Calibration -- **FeedbackCollector** — Records (confidence, action, was_correct) from human reviews -- **CalibrationCurve** — 5-bucket reliability analysis with ECE -- **IsotonicCalibrator** — Pure Python PAV algorithm for confidence calibration -- **ThresholdAdvisor** — Suggests optimal thresholds based on calibration data -- **ReliabilityDiagram** — ASCII art visualization for terminal - -### Evaluation & Experimentation -- **NLI Scorer** — Sentence decomposition, synonym expansion, negation detection -- **Retrieval Metrics** — Precision@K, Recall@K, MRR, NDCG -- **A/B Experiment Framework** — Same tickets, two configs, comparison report -- **Human Review Accuracy** — Precision/recall/F1 for review trigger correctness - -### Dashboard & Visualization -- **Confidence Dashboard** — Streamlit visualization of confidence distribution and tier routing -- **Retrieval Visualization** — Streamlit table + contribution chart for retrieval traces -- **Human Review Console** — Review interface with approve/edit/escalate/reject actions -- **Chat UI** — Multi-turn conversation interface with evidence panel and risk escalation - -### Observability -- **AgentTrace** — Append-only event stream per run -- **ClaimProvenance** — Answer → citation → chunk → document traceability -- **Provider Latency Measurement** — Benchmark script for LLM provider comparison - -## Quick Start - -### One-Click Demo - -```bash -# Check Docker, start DB, seed data, run demo, optionally launch dashboard -bash scripts/demo.sh -``` +大多数 AI 客服 demo 回避了最难的问题:**你怎么知道 LLM 没有在胡说?哪些工单需要人来判断?错误的回答怎么追溯到源头?** -### Manual Setup +TicketPilot 用工程手段回答这些问题: -```bash -git clone https://github.com/lennney/ticketpilot.git -cd ticketpilot +- **管线内零 LLM 调用** — 分类、风险、检索、评分全部确定性执行,结果可复现 +- **混合检索而非纯向量** — 关键词 FTS + pgvector HNSW → RRF 融合 → 4 信号混合重排序 +- **4 层置信度路由** — HIGH/MEDIUM 自动发送,LOW 人工审核,CRITICAL 强制转人工 +- **8 类禁止承诺检测** — 退款金额、法律威胁、隐私承诺等,AI 草稿不会越线 -pip install uv -uv sync +> Portfolio demo project. All data is synthetic. -cp .env.example .env.local -# Edit .env.local with your API keys (optional — pipeline works without LLM keys) +--- -docker compose up -d db +## 截图 -uv run python scripts/ingest_knowledge.py +| 监控大盘 | 置信度分布 | 意图×风险热力图 | +|:---:|:---:|:---:| +| ![Dashboard](docs/assets/dashboard-overview.png) | ![Charts](docs/assets/dashboard-charts.png) | ![Heatmap](docs/assets/dashboard-heatmap.png) | -uv run uvicorn ticketpilot.api:app --host 0.0.0.0 --port 8000 -``` +--- -### Run Tests +## 30 秒上手 ```bash -# Unit tests (no database required) -TICKETPILOT_SKIP_DB_TESTS=1 uv run pytest tests/ --ignore=tests/integration -q +git clone https://github.com/lennney/ticketpilot.git && cd ticketpilot -# Full quality gate -bash scripts/run_quality_gate.sh +pip install uv && uv sync # 安装依赖 +docker compose up -d db # 启动 PostgreSQL + pgvector +uv run python scripts/ingest_knowledge.py # 灌入知识库 + +uv run uvicorn ticketpilot.api:app --port 8000 # 启动 API ``` -### Review Console +```bash +# 一键 demo(灌数据 + 启动服务 + 跑评测) +bash scripts/demo.sh +``` ```bash +# 人工审核台 uv run streamlit run src/ticketpilot/review/console.py --server.port 8501 ``` -### Dashboard +--- -```bash -uv run python scripts/run_dashboard.py +## 架构 + +```mermaid +graph TD + A[工单输入] --> B[意图分类
8 类 + 置信度] + B --> C[风险评估
8 标记 × 3 级别] + C --> D[混合检索
FTS + pgvector → RRF] + D --> E[多 Agent 路由] + E --> F[退款 / 投诉 / 物流 / 技术 / 默认] + F --> G[草稿生成] + G --> H[Claim Guard
禁止承诺检测] + H --> I[置信度评分
4 维加权] + I --> J{路由决策} + J -->|HIGH / MEDIUM| K[自动发送] + J -->|LOW| L[人工审核] + J -->|CRITICAL| M[强制转人工] + L -->|通过| K + K --> N[反馈回路] + N --> O[等距校准器] + O --> B ``` -### Calibration & Feedback +--- + +## 和普通 RAG 的区别 + +| 维度 | 典型 RAG | TicketPilot | +|------|---------|-------------| +| 检索 | 单路向量搜索 | 关键词 FTS + 向量 HNSW → RRF → **4 信号混合重排序** | +| 置信度 | 二元(自信/不自信) | 4 维加权:检索 35% + 分类 25% + 引用 25% + 证据密度 15% | +| 路由 | 全自动或全人工 | 4 层降级:AUTO → CAUTIOUS → HUMAN_REVIEW → ESCALATION | +| 幻觉防护 | 无 | 8 类禁止承诺检测(退款金额、法律威胁等) | +| 可追溯性 | 无 | 全链路:回答 → 引用 → chunk → 文档 | +| Agent 架构 | 单 Agent | 5 个专职 Agent + 意图路由 | +| 管线确定性 | 依赖 LLM | 规则驱动,管线内零 LLM 调用 | +| 校准 | 静态阈值 | 反馈回路 + 等距回归 + 可靠性图 | + +--- + +## API + +| 端点 | 方法 | 说明 | +|------|------|------| +| `/api/tickets` | POST | 提交工单处理 | +| `/api/chat` | POST | 对话式 Copilot | +| `/api/chat/stream` | POST | SSE 流式响应 | +| `/api/reviews` | POST | 提交人工审核决策 | +| `/api/evaluation` | GET | 评测指标 | + +--- + +## 测试 ```bash -# Run calibration with reflection data -uv run python scripts/calibrate_with_reflection.py +# 单元测试(无需数据库) +TICKETPILOT_SKIP_DB_TESTS=1 uv run pytest tests/ --ignore=tests/integration -q -# Run A/B threshold experiment -uv run python scripts/run_threshold_ab.py +# 完整质量门禁(lint + 测试 + 集成 + openspec + 密钥扫描) +bash scripts/run_quality_gate.sh ``` -## API Endpoints +``` +1,760 tests passing · 87% coverage · ≥ 70% enforced +``` -| Endpoint | Method | Description | -|----------|--------|-------------| -| `/api/health` | GET | Health check | -| `/api/chat` | POST | Chat with AI copilot | -| `/api/chat/stream` | POST | Streaming chat (SSE) | -| `/api/tickets` | POST | Process ticket | -| `/api/reviews` | POST | Submit review decision | -| `/api/evaluation` | GET | Get evaluation metrics | +--- -## Project Structure +## 项目结构 ``` src/ticketpilot/ -├── api/ # FastAPI endpoints + SSE streaming -├── classification/ # Intent classifier (deterministic, 8 classes) -├── config/ # Central confidence thresholds -├── confidence/ # 4-dimensional confidence scorer -├── degradation/ # 4-tier response router -├── drafting/ # DraftAgent, prompt builder, claim guard, citation validator -├── evaluation/ # RAGAS-style metrics, NLI scorer, retrieval metrics, A/B experiments -├── experiment/ # A/B experiment framework (Config + Runner + Reporter) -├── feedback/ # Feedback collector, calibrator, threshold advisor -├── guardrails/ # PII detection, security scanning -├── intake/ # Ticket normalization, entity extraction -├── multi_agent/ # Orchestrator + 5 specialized agents -├── prompts/ # Per-agent prompt templates -├── retrieval/ # Hybrid retrieval (FTS + HNSW → RRF) -├── review/ # Streamlit review console, retrieval visualization -├── risk/ # Risk assessor + rules (8 flag types, 3 severity) -├── schema/ # Pydantic data models -├── tracing/ # Provenance tracking -└── triggers/ # CLI + webhook entry points +├── api/ # FastAPI + SSE 流式 +├── classification/ # 意图分类(确定性,8 类) +├── confidence/ # 4 维置信度评分 +├── degradation/ # 4 层响应路由 +├── drafting/ # 草稿生成 + Claim Guard + 引用验证 +├── evaluation/ # NLI 评分、检索指标、A/B 实验 +├── feedback/ # 反馈收集、等距校准、阈值顾问 +├── guardrails/ # PII 检测、安全扫描 +├── multi_agent/ # 编排器 + 5 专职 Agent +├── retrieval/ # 混合检索(FTS + HNSW → RRF → 混合重排序) +│ ├── hybrid_reranker.py # 多信号加权重排序 +│ ├── query_expander.py # LLM 查询扩展 +│ ├── result_merger.py # 多变体结果合并 +│ └── reranker_config.py # YAML 可配置权重 +├── review/ # Streamlit 人工审核台 +├── risk/ # 风险评估(8 标记,3 级别) +├── schema/ # Pydantic 数据模型 +└── tracing/ # 来源追溯 ``` -## Test Coverage +--- +## 参与贡献 + +欢迎 PR!适合入门的方向: + +- 📝 **文档** — 中英文使用示例、架构说明 +- 🧪 **测试** — 检索/分类的边界 case +- 🔧 **Bug 修复** — 看 [Issues](https://github.com/lennney/ticketpilot/issues) +- 🌐 **国际化** — 审核台多语言支持 + +```bash +uv sync --group dev # 安装开发依赖 +uv run pytest tests/ -v +bash scripts/run_quality_gate.sh # 提 PR 前跑一下 ``` -1,662 tests passing -├── Unit tests (no DB): 1,662 -├── Integration tests (DB required): separate -└── Coverage: 87% (>= 70% enforced) -``` -## Portfolio +--- + +## 技术文档 + +- [检索架构](docs/technical/retrieval_architecture.md) — 混合检索管线详解 +- [质量门禁](docs/technical/quality_gate.md) — 测试和验证规则 +- [项目 Portfolio](docs/portfolio/index.md) — 指标和 elevator pitch -See [docs/portfolio/index.md](docs/portfolio/index.md) for the project elevator pitch, architecture diagram, and key metrics. +--- ## License diff --git a/config/reranker.yaml b/config/reranker.yaml new file mode 100644 index 0000000..43f5e4e --- /dev/null +++ b/config/reranker.yaml @@ -0,0 +1,40 @@ +# Hybrid Reranker Configuration +# Weights must sum to 1.0 + +weights: + rrf_score: 0.40 + embedding_similarity: 0.25 + intent_metadata_boost: 0.20 + content_quality: 0.15 + +# Intent -> doc_type boost values +# Higher value = stronger preference for that doc_type given the intent +intent_boost: + refund: + policy: 0.15 + faq: 0.10 + return_exchange: + policy: 0.15 + faq: 0.10 + account_issue: + policy: 0.10 + faq: 0.10 + technical_issue: + faq: 0.10 + case: 0.10 + product_consulting: + faq: 0.15 + logistics: + faq: 0.10 + case: 0.10 + complaint: + case: 0.15 + policy: 0.10 + other: {} + +content_quality: + optimal_length_min: 200 + optimal_length_max: 800 + keyword_density_weight: 0.5 + +num_query_variants: 2 diff --git a/docs/assets/dashboard-charts.png b/docs/assets/dashboard-charts.png new file mode 100644 index 0000000..1ff3e56 Binary files /dev/null and b/docs/assets/dashboard-charts.png differ diff --git a/docs/assets/dashboard-heatmap.png b/docs/assets/dashboard-heatmap.png new file mode 100644 index 0000000..9d2a2c0 Binary files /dev/null and b/docs/assets/dashboard-heatmap.png differ diff --git a/docs/assets/dashboard-overview.png b/docs/assets/dashboard-overview.png new file mode 100644 index 0000000..6a6c295 Binary files /dev/null and b/docs/assets/dashboard-overview.png differ diff --git a/docs/technical/retrieval_architecture.md b/docs/technical/retrieval_architecture.md index d56c545..82a1c32 100644 --- a/docs/technical/retrieval_architecture.md +++ b/docs/technical/retrieval_architecture.md @@ -150,15 +150,79 @@ What the fake embedding provider **cannot** provide: - Meaningful ranking of results by relevance - Any real-world retrieval precision or recall +## Hybrid Reranking (Post-RRF) + +After RRF fusion, an optional **Hybrid Reranker** applies multi-signal weighted fusion +to improve Top-K ranking quality. This replaces the previous simple embedding tiebreaker. + +**Source:** `src/ticketpilot/retrieval/hybrid_reranker.py` + +### Signals + +| Signal | Default Weight | Description | +|--------|---------------|-------------| +| RRF score | 0.40 | Min-max normalized RRF fusion score | +| Embedding similarity | 0.25 | Cosine similarity (auto-disabled with FakeEmbedding) | +| Intent metadata boost | 0.20 | Intent → doc_type matching (e.g., refund→policy +0.15) | +| Content quality | 0.15 | Length appropriateness + keyword density | + +When FakeEmbeddingProvider is detected, the embedding signal weight is redistributed +proportionally among the other 3 signals (so weights always sum to 1.0). + +### Configuration + +Weights and parameters are configurable via `config/reranker.yaml`: + +```yaml +weights: + rrf_score: 0.40 + embedding_similarity: 0.25 + intent_metadata_boost: 0.20 + content_quality: 0.15 +``` + +**Source:** `src/ticketpilot/retrieval/reranker_config.py` + +## Multi-Query Expansion + +An optional **MultiQueryExpander** generates query variants using LLM to improve recall. +When enabled, the original query + N variants are each run through the retrieval pipeline +independently, then merged before hybrid reranking. + +**Source:** `src/ticketpilot/retrieval/query_expander.py` + +### Merge Strategies + +Results from multiple query variants are merged using: + +- **sum_score** (default): RRF scores summed per chunk_id — docs found by multiple + variants get higher scores (multi-path validation) +- **max_score**: Keep highest RRF score per chunk_id +- **rrf_again**: Apply second-level RRF across variant rankings + +**Source:** `src/ticketpilot/retrieval/result_merger.py` + +### Pipeline Integration + +``` +Query → (optional) MultiQueryExpander → N variants + → per-variant: keyword + vector + RRF + → merge (sum_score) + → HybridReranker (4-signal fusion) + → Top-K output + RetrievalTrace +``` + +Both features are backward-compatible: `enable_query_expansion=False` by default, +and `intent=None` disables intent boost. All new fields in RetrievalTrace are Optional. + ## Deferred Items The following retrieval refinements are explicitly deferred: -- **Real embedding provider** — Two tiers planned: small (384-d) and quality (768-d). The `EmbeddingProvider` protocol interface is ready for integration. - **Realistic enterprise data pack** — Current 36-document seed set is synthetic. A real data pack with actual FAQ, policy, and case documents is needed. - **SourceRouter implementation** — Intent-to-source routing (e.g., refund tickets search only FAQ + Policy) was designed but not implemented. - **Persistent retrieval traces** — `RetrievalTrace` is in-memory only. The `retrieval_traces` DB table migration is deferred. -- **Retrieval evaluation** — No golden question-answer pairs, no precision/recall/mRR metrics, no evaluation harness. - **BM25 or alternative keyword retrieval** — PostgreSQL FTS is sufficient for MVP. BM25 may improve keyword ranking. - **Embedding fine-tuning** — No support ticket data available for fine-tuning. - **Evidence scoring threshold tuning** — RRF scores have no absolute meaning; threshold tuning deferred until evaluation data exists. +- **Cross-encoder reranker** — Would require sentence-transformers dependency; deferred in favor of lightweight multi-signal fusion. diff --git a/openspec/changes/add-hybrid-retrieval-reranking/design.md b/openspec/changes/add-hybrid-retrieval-reranking/design.md new file mode 100644 index 0000000..2ec515d --- /dev/null +++ b/openspec/changes/add-hybrid-retrieval-reranking/design.md @@ -0,0 +1,306 @@ +# Design: Hybrid Retrieval Reranking + +## Architecture Overview + +``` + ┌─────────────────────────────────────┐ + │ MultiQueryExpander │ + │ (LLM 生成 2 个查询变体 + 原始查询) │ + └──────────────┬──────────────────────┘ + │ 3 queries + ┌──────────────▼──────────────────────┐ + │ Parallel Retrieval (3 路) │ + │ keyword_search + vector_search │ + │ per query → RRF fusion │ + └──────────────┬──────────────────────┘ + │ 3 × top_2k fused results + ┌──────────────▼──────────────────────┐ + │ Result Merger + Dedup │ + │ (chunk_id 去重, 取最高 RRF score) │ + └──────────────┬──────────────────────┘ + │ merged candidates + ┌──────────────▼──────────────────────┐ + │ HybridReranker │ + │ │ + │ signal_1: rrf_score (w1) │ + │ signal_2: embedding_sim (w2) │ + │ signal_3: intent_meta_boost (w3) │ + │ signal_4: content_quality (w4) │ + │ │ + │ final_score = Σ(wi × normalized_i) │ + └──────────────┬──────────────────────┘ + │ reranked top_k + ┌──────────────▼──────────────────────┐ + │ RetrievalTrace │ + │ (所有信号 + 权重 + 最终分数) │ + └─────────────────────────────────────┘ +``` + +## Component Design + +### 1. MultiQueryExpander + +**File**: `src/ticketpilot/retrieval/query_expander.py` + +```python +class MultiQueryExpander: + """Generate query variants using LLM for improved recall.""" + + def __init__(self, llm_client=None, num_variants: int = 2): + self._llm = llm_client # Reuse DraftAgent's LLM config + self._num_variants = num_variants + + def expand(self, query: str, intent: str = "") -> list[str]: + """Return [original_query, variant_1, variant_2, ...]. + + On LLM failure, returns [original_query] only. + """ +``` + +**LLM Prompt**: +``` +你是一个搜索查询优化器。给定一个客服工单查询,生成 {n} 个不同角度的搜索关键词变体。 +要求:每个变体 5-15 个字,覆盖不同语义角度(同义词、上位词、具体化)。 +只输出 JSON 数组,不要解释。 + +查询:{query} +意图:{intent} +``` + +**输出**: `["退款到账时间", "退款进度查询"]` + +**Fallback chain**: +1. LLM 成功 → 返回变体 +2. LLM 超时/报错 → 返回 `[original_query]` + 日志警告 +3. 无 API key → 跳过扩展,返回 `[original_query]` + +### 2. ResultMerger + +**File**: `src/ticketpilot/retrieval/result_merger.py` + +```python +def merge_retrieval_results( + result_sets: list[list[FusedResult]], + strategy: str = "max_score", # or "sum_score", "rrf_again" +) -> list[FusedResult]: + """Merge multiple retrieval result sets, deduplicating by chunk_id. + + strategy="max_score": keep highest RRF score per chunk_id + strategy="sum_score": sum RRF scores across queries (boosts docs found by multiple queries) + strategy="rrf_again": treat each query as a ranker, apply second-level RRF + """ +``` + +**推荐策略**: `sum_score` — 被多个查询变体命中的文档得分更高(类似 PageIndex 的多路径验证思想)。 + +### 3. HybridReranker + +**File**: `src/ticketpilot/retrieval/hybrid_reranker.py`(替代现有 `reranker.py`) + +```python +@dataclass +class RerankSignal: + """One scoring signal with its weight and raw/normalized values.""" + name: str + weight: float + raw_value: float + normalized_value: float + contribution: float # weight * normalized_value + +@dataclass +class RerankResult: + """Reranked result with signal breakdown.""" + chunk_id: UUID + final_score: float + signals: list[RerankSignal] + rank: int + +class HybridReranker: + """Multi-signal reranker combining RRF, embedding, intent, and content signals.""" + + def __init__(self, config: RerankerConfig | None = None): + self._config = config or RerankerConfig.default() + + def rerank( + self, + candidates: list[FusedResult], + query: str, + query_embedding: list[float] | None, + intent: IntentClass | None, + top_k: int = 10, + ) -> list[RerankResult]: + """Rerank candidates using weighted multi-signal fusion.""" +``` + +#### Signal 1: RRF Score (weight: 0.4) +- 直接使用现有 RRF score +- Min-max normalization: `(score - min) / (max - min)` + +#### Signal 2: Embedding Similarity (weight: 0.25) +- cosine_similarity(query_embedding, doc_embedding) +- 需要真实 embedding 才有意义 +- FakeEmbedding 时此信号权重自动降为 0,重新分配 + +#### Signal 3: Intent Metadata Boost (weight: 0.2) +- 基于 `IntentClass` → `doc_type` 匹配表 +- 匹配: +1.0, 不匹配: 0.0 +- 二值信号,不做连续打分 + +#### Signal 4: Content Quality (weight: 0.15) +- `length_score`: 内容长度适中(200-800字)得分最高,太短/太长扣分 +- `keyword_density`: 查询关键词在内容中的命中比例 +- 两个子信号取平均 + +#### Weight Auto-adjustment +```python +def _adjust_weights(self, has_real_embedding: bool) -> dict[str, float]: + """Redistribute weights when signals are unavailable.""" + weights = self._config.weights.copy() + if not has_real_embedding: + # Remove embedding signal, redistribute proportionally + embedding_weight = weights.pop("embedding_similarity") + total = sum(weights.values()) + weights = {k: v / total * 1.0 for k, v in weights.items()} + return weights +``` + +### 4. RerankerConfig + +**File**: `src/ticketpilot/retrieval/reranker_config.py` + +```python +@dataclass +class RerankerConfig: + weights: dict[str, float] # signal_name -> weight + intent_boost_table: dict[str, dict[str, float]] # intent -> {doc_type: boost} + content_quality: ContentQualityConfig + enable_llm_scoring: bool = False # Phase 2: LLM-based relevance + + @classmethod + def default(cls) -> "RerankerConfig": + """Default config with balanced weights.""" + + @classmethod + def from_yaml(cls, path: str) -> "RerankerConfig": + """Load from config file for A/B experiments.""" +``` + +**Config file**: `config/reranker.yaml` +```yaml +weights: + rrf_score: 0.40 + embedding_similarity: 0.25 + intent_metadata_boost: 0.20 + content_quality: 0.15 + +intent_boost: + refund: + policy: 0.15 + faq: 0.10 + complaint: + case: 0.15 + policy: 0.10 + # ... + +content_quality: + optimal_length_min: 200 + optimal_length_max: 800 + keyword_density_weight: 0.5 +``` + +### 5. RetrievalTrace 扩展 + +现有 `RetrievalTrace` 新增字段: + +```python +@dataclass +class RetrievalTrace: + # ... existing fields ... + + # New fields for hybrid reranking + query_variants: list[str] | None = None # 扩展查询列表 + expansion_latency_ms: int = 0 + merged_result_count: int = 0 # 去重后候选数 + rerank_signals: list[dict] | None = None # 每个结果的信号分解 + reranker_weights: dict[str, float] | None = None # 实际使用的权重 + has_real_embedding: bool = False # 是否使用真实 embedding +``` + +### 6. Pipeline 集成 + +修改 `src/ticketpilot/retrieval/pipeline.py`: + +```python +def hybrid_retrieval( + query: str, + top_k: int = 10, + intent: IntentClass | None = None, # NEW: 用于 intent boost + # ... existing params ... + enable_query_expansion: bool = True, # NEW: 多查询扩展开关 + reranker_config: RerankerConfig | None = None, # NEW: 重排配置 +) -> RetrievalTrace: + """Enhanced hybrid retrieval with multi-query expansion and hybrid reranking.""" + + # Step 0: Query expansion (optional) + if enable_query_expansion: + expander = MultiQueryExpander() + queries = expander.expand(query, intent.value if intent else "") + else: + queries = [query] + + # Step 1-3: Parallel retrieval per query + all_fused = [] + for q in queries: + trace_q = _single_query_retrieval(q, top_k, doc_types, ...) + all_fused.append(trace_q.fused_results) + + # Step 4: Merge + dedup + merged = merge_retrieval_results(all_fused, strategy="sum_score") + + # Step 5: Hybrid rerank + reranker = HybridReranker(config=reranker_config) + reranked = reranker.rerank( + candidates=merged, + query=query, + query_embedding=query_embedding, + intent=intent, + top_k=top_k, + ) + + # Step 6: Build trace + return RetrievalTrace(...) +``` + +## File Manifest + +| File | Action | Description | +|------|--------|-------------| +| `src/ticketpilot/retrieval/query_expander.py` | NEW | MultiQueryExpander | +| `src/ticketpilot/retrieval/result_merger.py` | NEW | Result merge + dedup | +| `src/ticketpilot/retrieval/hybrid_reranker.py` | NEW | 多信号混合重排器 | +| `src/ticketpilot/retrieval/reranker_config.py` | NEW | 重排配置 dataclass | +| `config/reranker.yaml` | NEW | 默认权重配置 | +| `src/ticketpilot/retrieval/pipeline.py` | MODIFY | 集成 expansion + hybrid rerank | +| `src/ticketpilot/retrieval/traces.py` | MODIFY | 新增 trace 字段 | +| `src/ticketpilot/retrieval/retrieve_evidence.py` | MODIFY | 传递 intent 参数 | +| `tests/unit/test_query_expander.py` | NEW | 查询扩展单元测试 | +| `tests/unit/test_result_merger.py` | NEW | 结果合并单元测试 | +| `tests/unit/test_hybrid_reranker.py` | NEW | 混合重排单元测试 | +| `tests/unit/test_pipeline_retrieval.py` | MODIFY | 更新 pipeline 测试 | +| `reports/retrieval/hybrid_rerank_comparison.md` | NEW | before/after 对比报告 | + +## Compatibility + +- `retrieve_evidence()` 接口向后兼容:新增 `intent` 参数,可选 +- `hybrid_retrieval()` 接口向后兼容:新增参数均有默认值 +- FakeEmbeddingProvider 继续作为无网络 fallback +- 现有 RetrievalTrace 消费者(dashboard, evaluation)不受影响(新字段都是 Optional) +- `reranker.py` 保留不删除(`rerank_with_embeddings` 和 `rerank_with_cross_encoder`),新 reranker 是独立模块 + +## Safety Constraints + +- 多查询扩展 LLM 调用失败必须 graceful fallback(返回原始查询) +- HybridReranker 权重总和必须 = 1.0(运行时校验) +- 无真实 embedding 时自动降级(embedding 信号权重归零重新分配) +- 所有配置通过 YAML 文件管理,不硬编码 +- 质量门必须通过 diff --git a/openspec/changes/add-hybrid-retrieval-reranking/proposal.md b/openspec/changes/add-hybrid-retrieval-reranking/proposal.md new file mode 100644 index 0000000..3507212 --- /dev/null +++ b/openspec/changes/add-hybrid-retrieval-reranking/proposal.md @@ -0,0 +1,118 @@ +# Proposal: Hybrid Retrieval Reranking + +## Executive Summary + +TicketPilot 当前检索管线:keyword FTS + pgvector HNSW → RRF fusion → embedding tiebreaker rerank。Phase 8 已实现 OpenAICompatibleProvider 但默认仍是 FakeEmbeddingProvider,reranker 只用 embedding 相似度做 tiebreaker,cross-encoder 是空 TODO。 + +本次改造引入**混合重排器(Hybrid Reranker)**:在 RRF fusion 之后,用多信号加权融合替代单一 embedding tiebreaker,显著提升 Top-K 排序质量。同时加入**多查询扩展**提升召回率。 + +灵感来源:PageIndex 的 LLM 树搜索思路,但适配 TicketPilot 的客服场景——不需要建树(文档短),而是用 LLM 做查询扩展和相关性评分。 + +## Baseline (Current State) + +### Retrieval Pipeline +``` +Query → build_retrieval_query (静态意图词映射, _INTENT_TERMS 硬编码) + → keyword_search (PostgreSQL FTS 'simple' + 32个业务词 LIKE 兜底) + → vector_search (pgvector HNSW, FakeEmbedding 384-dim) + → RRF fusion (k=60) + → rerank_with_embeddings (embedding相似度作tiebreaker, 无实际语义) + → top_k output +``` + +### Known Gaps +1. **FakeEmbedding 无语义**: 向量搜索路输出随机排序,RRF 有一半输入是噪声 +2. **静态查询扩展**: `_INTENT_TERMS` 硬编码,"退款"↔"退钱"↔"返还费用" 无法覆盖 +3. **Reranker 弱**: 仅 embedding tiebreaker,cross-encoder 是空 TODO +4. **无意图感知排序**: 退款工单检索到物流文档不会被降权 +5. **无查询多样性**: 单一查询表达,语义变体覆盖不足 + +## Goal + +1. 实现 **HybridReranker**:多信号加权融合(RRF score + embedding similarity + intent metadata boost + content quality signal) +2. 实现 **MultiQueryExpander**:基于 LLM 生成 2-3 个查询变体,并行检索后合并去重 +3. 接入真实 embedding 作为默认(保留 FakeEmbedding 用于无网络测试) +4. 所有新信号记录到 RetrievalTrace,支持调试和评估 +5. 用现有 101 eval tickets 做 before/after 对比 + +## Non-goals + +- ❌ 不做 PageIndex 树搜索(TicketPilot 文档平均 500-1000 字,无需层级目录) +- ❌ 不做 cross-encoder reranker(需要 sentence-transformers 重依赖,与 uv 轻量原则冲突) +- ❌ 不改知识库 schema(不加文档摘要字段,留到下个迭代) +- ❌ 不改 DraftAgent 内部逻辑 +- ❌ 不做 embedding fine-tuning +- ❌ 不做生产部署 +- ❌ 不 commit API key + +## Key Design Decisions + +### A. Hybrid Reranker 信号融合 + +| 信号 | 权重范围 | 来源 | 说明 | +|------|----------|------|------| +| RRF score | 0.3-0.5 | 现有 | keyword + vector 融合排名 | +| Embedding similarity | 0.2-0.3 | 现有(需真实embedding) | query-document 语义相似度 | +| Intent metadata boost | 0.1-0.2 | 新增 | 意图分类→文档类型匹配加分 | +| Content quality signal | 0.05-0.1 | 新增 | 内容长度、关键词密度等启发式 | + +权重通过配置文件管理,支持 A/B 实验。 + +### B. Intent Metadata Boost 逻辑 + +| IntentClass | 优先 doc_type | 加分 | +|-------------|--------------|------| +| REFUND | policy, faq | +0.15 | +| RETURN_EXCHANGE | policy, faq | +0.15 | +| COMPLAINT | case, policy | +0.15 | +| TECHNICAL_ISSUE | faq, case | +0.1 | +| ACCOUNT_ISSUE | policy, faq | +0.1 | +| LOGISTICS | faq, case | +0.1 | +| PRODUCT_CONSULTING | faq | +0.15 | +| OTHER | (无加分) | 0 | + +### C. Multi-Query Expansion + +``` +原始查询: "我买的东西退款一直没到账" + ↓ LLM expansion (DeepSeek-chat, temperature=0.5) +扩展查询: ["退款到账时间", "退款进度查询", "退款未收到怎么办"] + ↓ 并行检索 (3路) + ↓ RRF merge + dedup + ↓ HybridReranker +最终 Top-K +``` + +- 每次扩展生成 2 个变体(控制 token 消耗) +- DeepSeek-chat 调用,单次 ~200 tokens,成本可忽略 +- 扩展失败时 graceful fallback 到原始查询 + +### D. 真实 Embedding 策略 + +| 决策 | 选择 | +|------|------| +| 默认 provider | `EMBEDDING_PROVIDER` env var 控制,默认 `openai_compatible` | +| 无网络 fallback | 自动降级到 FakeEmbeddingProvider + 日志警告 | +| 模型 | `text-embedding-3-small` (1536-dim) 或 DashScope `text-embedding-v4` (1024-dim) | +| 索引重建 | 维度变更时自动检测 + 提示重建 | + +## Proposed Metrics + +| Metric | Definition | Baseline Target | +|--------|-----------|----------------| +| Top-3 hit rate | Top-3 检索命中 golden expected doc | 提升 10%+ | +| Top-5 hit rate | Top-5 检索命中 | 提升 8%+ | +| MRR | Mean Reciprocal Rank | 提升 15%+ | +| Intent-aware precision | 检索结果 doc_type 与意图匹配率 | 新指标 | +| Reranker latency | 单次重排耗时 | < 50ms (无LLM) / < 500ms (含LLM) | +| Query expansion coverage | 扩展查询召回原始查询未命中的文档比例 | 新指标 | + +## Constraints + +- FakeEmbeddingProvider 必须保留为无网络环境的 fallback +- HybridReranker 权重必须可配置(支持 A/B) +- 多查询扩展的 LLM 调用失败必须 graceful fallback +- 所有新信号必须记录到 RetrievalTrace +- 现有 101 eval tickets 不修改 +- 质量门必须通过(ruff + tests + openspec) +- 无 API key commit diff --git a/openspec/changes/add-hybrid-retrieval-reranking/specs/hybrid-reranking/spec.md b/openspec/changes/add-hybrid-retrieval-reranking/specs/hybrid-reranking/spec.md new file mode 100644 index 0000000..b5bc101 --- /dev/null +++ b/openspec/changes/add-hybrid-retrieval-reranking/specs/hybrid-reranking/spec.md @@ -0,0 +1,159 @@ +# hybrid-reranking Specification + +## Purpose +Define the hybrid reranking system that combines multiple scoring signals (RRF score, embedding similarity, intent metadata boost, content quality) to produce higher-quality Top-K retrieval results for customer support ticket evidence retrieval. + +## Requirements + +### Requirement: HybridReranker multi-signal fusion +The system SHALL implement a HybridReranker that combines at least 4 scoring signals with configurable weights. + +#### Scenario: HybridReranker produces ranked results +- **WHEN** HybridReranker.rerank() is called with candidates, query, and config +- **THEN** returns a list of RerankResult sorted by final_score descending + +#### Scenario: Signal weights sum to 1.0 +- **WHEN** RerankerConfig is loaded +- **THEN** all signal weights sum to 1.0 (validated at load time) + +#### Scenario: Weight auto-adjustment on missing signals +- **WHEN** a signal is unavailable (e.g., fake embedding) +- **THEN** its weight is redistributed proportionally among available signals + +### Requirement: RRF Score Signal +The system SHALL use the existing RRF fusion score as the primary reranking signal. + +#### Scenario: RRF score normalization +- **WHEN** RRF scores are processed +- **THEN** scores are min-max normalized to [0, 1] range within the candidate set + +### Requirement: Embedding Similarity Signal +The system SHALL compute cosine similarity between query embedding and document embedding as a reranking signal. + +#### Scenario: Real embedding similarity +- **WHEN** real embedding provider is active +- **THEN** embedding_similarity signal contributes to final score per configured weight + +#### Scenario: Fake embedding auto-downgrade +- **WHEN** FakeEmbeddingProvider is detected +- **THEN** embedding_similarity weight is set to 0 and redistributed to other signals + +### Requirement: Intent Metadata Boost Signal +The system SHALL boost documents whose doc_type matches the classified intent. + +#### Scenario: Intent-doc_type match +- **WHEN** intent is REFUND and doc_type is "policy" +- **THEN** intent_metadata_boost score = 1.0 (boost applied) + +#### Scenario: Intent-doc_type mismatch +- **WHEN** intent is REFUND and doc_type is "case" +- **THEN** intent_metadata_boost score = 0.0 (no boost) + +#### Scenario: No intent available +- **WHEN** intent is None +- **THEN** intent_metadata_boost score = 0.0 for all candidates + +### Requirement: Content Quality Signal +The system SHALL score documents based on content length appropriateness and keyword density. + +#### Scenario: Optimal length content scores highest +- **WHEN** content length is between 200-800 characters +- **THEN** length_score is at its peak + +#### Scenario: Very short content scores lower +- **WHEN** content length < 50 characters +- **THEN** length_score is significantly below peak + +#### Scenario: Keyword density scoring +- **WHEN** query contains "退款" and document contains "退款" 3 times +- **THEN** keyword_density is higher than a document containing it 0 times + +### Requirement: MultiQueryExpander +The system SHALL generate query variants using LLM to improve recall. + +#### Scenario: Successful expansion +- **WHEN** expand("退款没到账", intent="refund") is called with LLM available +- **THEN** returns ["退款没到账", variant_1, variant_2] (3 queries total) + +#### Scenario: LLM failure fallback +- **WHEN** LLM call fails (timeout, error, no API key) +- **THEN** returns ["退款没到账"] (original query only) with warning logged + +#### Scenario: Variant quality control +- **WHEN** LLM returns variants longer than 50 characters or empty +- **THEN** invalid variants are filtered out + +### Requirement: ResultMerger +The system SHALL merge results from multiple query retrievals with deduplication. + +#### Scenario: Sum-score merge +- **WHEN** chunk X appears in query_1 results (score=0.3) and query_2 results (score=0.2) +- **THEN** merged score for chunk X = 0.5 (sum) + +#### Scenario: Deduplication +- **WHEN** same chunk_id appears in multiple result sets +- **THEN** only one entry in merged results with aggregated score + +#### Scenario: Empty result sets +- **WHEN** all result sets are empty +- **THEN** returns empty list + +### Requirement: RetrievalTrace extension +The system SHALL record all hybrid reranking signals in the retrieval trace. + +#### Scenario: Trace records query variants +- **WHEN** query expansion is used +- **THEN** trace.query_variants contains all query strings used + +#### Scenario: Trace records per-result signals +- **WHEN** hybrid reranking completes +- **THEN** trace.rerank_signals contains signal breakdown for each result + +#### Scenario: Trace records actual weights used +- **WHEN** weights are auto-adjusted +- **THEN** trace.reranker_weights reflects the adjusted weights + +#### Scenario: Trace records embedding provider status +- **WHEN** pipeline completes +- **THEN** trace.has_real_embedding indicates if real embedding was used + +### Requirement: Pipeline backward compatibility +The system SHALL maintain backward compatibility with existing callers. + +#### Scenario: Existing retrieve_evidence call without intent +- **WHEN** retrieve_evidence() is called without intent parameter +- **THEN** works identically to before (intent=None, no intent boost) + +#### Scenario: Existing hybrid_retrieval call without new params +- **WHEN** hybrid_retrieval() is called without enable_query_expansion and reranker_config +- **THEN** uses defaults (expansion enabled, default reranker config) + +### Requirement: RerankerConfig from YAML +The system SHALL support loading reranker configuration from YAML files. + +#### Scenario: Load default config +- **WHEN** RerankerConfig.default() is called +- **THEN** returns config with balanced weights (0.40, 0.25, 0.20, 0.15) + +#### Scenario: Load from YAML file +- **WHEN** RerankerConfig.from_yaml("config/reranker.yaml") is called +- **THEN** loads and validates weights from file + +#### Scenario: Invalid YAML weights +- **WHEN** YAML weights sum to 0.8 (not 1.0) +- **THEN** raises ValueError with descriptive message + +### Requirement: Graceful degradation +The system SHALL degrade gracefully when components are unavailable. + +#### Scenario: No LLM for query expansion +- **WHEN** no LLM API key configured +- **THEN** query expansion is skipped, pipeline continues with original query + +#### Scenario: No real embedding provider +- **WHEN** FakeEmbeddingProvider is the only available provider +- **THEN** reranker runs with embedding signal weight = 0, other signals adjusted + +#### Scenario: RerankerConfig file missing +- **WHEN** config/reranker.yaml does not exist +- **THEN** falls back to RerankerConfig.default() diff --git a/openspec/changes/add-hybrid-retrieval-reranking/tasks.md b/openspec/changes/add-hybrid-retrieval-reranking/tasks.md new file mode 100644 index 0000000..e9a0a4b --- /dev/null +++ b/openspec/changes/add-hybrid-retrieval-reranking/tasks.md @@ -0,0 +1,111 @@ +# Tasks: Hybrid Retrieval Reranking + +## Phase 1: RerankerConfig + 配置文件 (30 min) + +### Task 1.1: Create `reranker_config.py` +- [ ] `RerankerConfig` dataclass: weights, intent_boost_table, content_quality +- [ ] `ContentQualityConfig` dataclass: optimal_length_min/max, keyword_density_weight +- [ ] `RerankerConfig.default()` class method +- [ ] `RerankerConfig.from_yaml(path)` class method +- [ ] `validate()` 方法:校验权重总和 = 1.0 +- [ ] Unit tests: 默认配置、YAML 加载、权重校验 + +### Task 1.2: Create `config/reranker.yaml` +- [ ] 默认权重: rrf_score=0.40, embedding_similarity=0.25, intent_metadata_boost=0.20, content_quality=0.15 +- [ ] Intent boost 表: 8 个 IntentClass × 优先 doc_type +- [ ] Content quality 参数: optimal_length_min=200, optimal_length_max=800 +- [ ] Unit test: YAML 可解析且权重合法 + +## Phase 2: HybridReranker 核心 (45 min) + +### Task 2.1: Create `hybrid_reranker.py` — Signal 1 (RRF Score) +- [ ] `RerankSignal` dataclass: name, weight, raw_value, normalized_value, contribution +- [ ] `RerankResult` dataclass: chunk_id, final_score, signals, rank +- [ ] `HybridReranker.rerank()` 骨架 +- [ ] Signal 1: RRF score min-max normalization +- [ ] Unit test: 单信号 rerank 结果 = RRF 排序 + +### Task 2.2: Signal 2 (Embedding Similarity) +- [ ] cosine_similarity 计算(复用 reranker.py 现有函数) +- [ ] 从 DB 获取 doc embedding(复用 `_get_document_embedding`) +- [ ] FakeEmbedding 检测:自动降级(权重归零重分配) +- [ ] Unit test: 真实 embedding 时相似度参与打分;fake 时自动降级 + +### Task 2.3: Signal 3 (Intent Metadata Boost) +- [ ] `intent_boost_table` 查询逻辑 +- [ ] IntentClass → doc_type 匹配 → 加分 +- [ ] Unit test: 退款意图 + policy 文档 → +0.15;退款意图 + logistics 文档 → 0 + +### Task 2.4: Signal 4 (Content Quality) +- [ ] `length_score`: 正态分布曲线,optimal_length 中心峰值最高 +- [ ] `keyword_density`: 查询词在 content 中的命中比例 +- [ ] 两个子信号取平均作为 content_quality signal +- [ ] Unit test: 短/中/长内容得分差异;高/低关键词密度得分差异 + +### Task 2.5: Weight Auto-adjustment + Integration +- [ ] `_adjust_weights()`: 信号不可用时重新分配权重 +- [ ] 4 信号加权求和 → final_score +- [ ] 结果按 final_score 降序排列,赋 rank +- [ ] Unit test: 权重总和始终 = 1.0;4 信号融合正确 + +## Phase 3: MultiQueryExpander (30 min) + +### Task 3.1: Create `query_expander.py` +- [ ] `MultiQueryExpander` class +- [ ] LLM prompt 模板(中文,JSON 输出) +- [ ] `expand()` → `[original, variant_1, variant_2]` +- [ ] JSON 解析 + 校验(长度、数量) +- [ ] Fallback: LLM 失败 → `[original]` + 日志警告 +- [ ] Unit test: 正常扩展、LLM 失败 fallback、无 API key 跳过 + +## Phase 4: ResultMerger (20 min) + +### Task 4.1: Create `result_merger.py` +- [ ] `merge_retrieval_results(result_sets, strategy="sum_score")` +- [ ] `sum_score`: 同一 chunk_id 在多路结果中得分求和 +- [ ] `max_score`: 取最高 RRF score +- [ ] `rrf_again`: 对多路排名做二次 RRF +- [ ] Dedup by chunk_id,保留最高分版本的 content +- [ ] Unit test: 3 路结果合并去重;sum_score 加分逻辑 + +## Phase 5: Pipeline 集成 (30 min) + +### Task 5.1: 修改 `traces.py` +- [ ] 新增字段: query_variants, expansion_latency_ms, merged_result_count +- [ ] 新增字段: rerank_signals, reranker_weights, has_real_embedding +- [ ] 所有新字段 Optional,默认 None/0/False +- [ ] 现有测试不受影响 + +### Task 5.2: 修改 `pipeline.py` +- [ ] `hybrid_retrieval()` 新增参数: intent, enable_query_expansion, reranker_config +- [ ] Step 0: query expansion (if enabled) +- [ ] Step 1-3: 并行检索 per query variant +- [ ] Step 4: merge_retrieval_results +- [ ] Step 5: HybridReranker.rerank (替代现有 rerank_with_embeddings) +- [ ] Step 6: 构建扩展 trace +- [ ] 向后兼容:新参数都有默认值,现有调用不受影响 + +### Task 5.3: 修改 `retrieve_evidence.py` +- [ ] 新增 `intent` 参数传递到 `hybrid_retrieval()` +- [ ] 向后兼容:intent 默认 None + +### Task 5.4: 更新现有测试 +- [ ] `test_pipeline_retrieval.py`: 新增 hybrid rerank 路径测试 +- [ ] `test_retrieve_evidence.py`: 新增 intent 传递测试 +- [ ] 所有现有测试必须继续通过 + +## Phase 6: Evaluation + Report (30 min) + +### Task 6.1: Before/After 对比 +- [ ] 用现有 101 eval tickets 跑 before(当前 pipeline) +- [ ] 跑 after(hybrid reranker) +- [ ] 对比指标: Top-3 hit rate, Top-5 hit rate, MRR +- [ ] 输出 `reports/retrieval/hybrid_rerank_comparison.md` + +### Task 6.2: 质量门 +- [ ] `ruff check` 通过 +- [ ] `pytest` 全部通过 +- [ ] `openspec validate --all` 通过 +- [ ] Secret scan 通过 + +## Total Estimated Time: ~3 hours diff --git a/src/ticketpilot/drafting/draft_agent.py b/src/ticketpilot/drafting/draft_agent.py index 51efc9a..42d07ba 100644 --- a/src/ticketpilot/drafting/draft_agent.py +++ b/src/ticketpilot/drafting/draft_agent.py @@ -39,6 +39,8 @@ # Minimum RRF score threshold for evidence to be considered "good" _EVIDENCE_SCORE_THRESHOLD = 0.01 +# Maximum evidence items to keep (prevents unbounded context growth) +_MAX_EVIDENCE = 15 # Maximum agent loop iterations (safety bound) _MAX_ITERATIONS = 5 # Safe fallback when agent cannot produce a grounded reply @@ -355,9 +357,10 @@ def generate_draft( } ) - # Seed state with any pre-retrieved evidence + # Seed state with any pre-retrieved evidence (capped to prevent context overflow) if evidence_candidates: - state.evidence = list(evidence_candidates) + sorted_candidates = sorted(evidence_candidates, key=lambda e: e.score, reverse=True) + state.evidence = sorted_candidates[:_MAX_EVIDENCE] try: result = self._run_agent_loop( @@ -790,9 +793,14 @@ def _reformulate_search( state.evidence.append(c) existing_ids.add(c.chunk_id) + # Cap evidence by score to prevent unbounded context growth + state.evidence.sort(key=lambda e: e.score, reverse=True) + state.evidence = state.evidence[:_MAX_EVIDENCE] + logger.info( - "DraftAgent: reformulated search added %d new results", + "DraftAgent: reformulated search added %d new results (total %d)", len(new_candidates), + len(state.evidence), ) def _llm_guided_search( @@ -832,9 +840,19 @@ def _llm_guided_search( if query and query not in state.search_queries_used: state.search_queries_used.append(query) raw_results = _search_knowledge(query) - state.evidence = self._raw_results_to_candidates(raw_results) + # Merge with existing evidence (don't replace — preserve good earlier results) + new_candidates = self._raw_results_to_candidates(raw_results) + existing_ids = {c.chunk_id for c in state.evidence} + for c in new_candidates: + if c.chunk_id not in existing_ids: + state.evidence.append(c) + existing_ids.add(c.chunk_id) + # Cap by score + state.evidence.sort(key=lambda e: e.score, reverse=True) + state.evidence = state.evidence[:_MAX_EVIDENCE] logger.info( - "DraftAgent: LLM-guided search returned %d results", + "DraftAgent: LLM-guided search added %d new results (total %d)", + len(new_candidates), len(state.evidence), ) except Exception as e: diff --git a/src/ticketpilot/retrieval/hybrid_reranker.py b/src/ticketpilot/retrieval/hybrid_reranker.py new file mode 100644 index 0000000..6f32bc4 --- /dev/null +++ b/src/ticketpilot/retrieval/hybrid_reranker.py @@ -0,0 +1,315 @@ +"""Hybrid reranker combining multiple scoring signals. + +Signals: +1. RRF score (from keyword + vector fusion) +2. Embedding similarity (cosine similarity, requires real embedding) +3. Intent metadata boost (intent -> doc_type matching) +4. Content quality (length appropriateness + keyword density) + +All signals are normalized to [0, 1] and combined with configurable weights. +""" +from __future__ import annotations + +import math +import re +import logging +from dataclasses import dataclass, field +from typing import Any, Optional +from uuid import UUID + +from ticketpilot.retrieval.reranker_config import RerankerConfig +from ticketpilot.retrieval.traces import FusedResult + +logger = logging.getLogger(__name__) + +# --------------------------------------------------------------------------- +# Output dataclasses +# --------------------------------------------------------------------------- + +@dataclass +class RerankSignal: + """One scoring signal with its weight and raw/normalized values.""" + name: str + weight: float + raw_value: float + normalized_value: float + contribution: float # weight * normalized_value + + +@dataclass +class RerankResult: + """Reranked result with signal breakdown.""" + chunk_id: UUID + doc_id: UUID + doc_type: str + content: str + final_score: float + signals: list[RerankSignal] = field(default_factory=list) + rank: int = 0 + # Preserve original RRF info + rrf_score: float = 0.0 + keyword_rank: Optional[int] = None + keyword_contribution: Optional[float] = None + vector_rank: Optional[int] = None + vector_contribution: Optional[float] = None + sources: list[str] = field(default_factory=list) + + def to_fused_result(self) -> FusedResult: + """Convert back to FusedResult for downstream compatibility.""" + from ticketpilot.retrieval.schema.knowledge import DocType # noqa: PLC0415 + return FusedResult( + chunk_id=self.chunk_id, + doc_id=self.doc_id, + doc_type=DocType(self.doc_type) if isinstance(self.doc_type, str) else self.doc_type, + content=self.content, + rrf_score=self.rrf_score, + keyword_rank=self.keyword_rank, + keyword_contribution=self.keyword_contribution, + vector_rank=self.vector_rank, + vector_contribution=self.vector_contribution, + sources=self.sources + ["hybrid_rerank"], + ) + + +# --------------------------------------------------------------------------- +# Signal computations +# --------------------------------------------------------------------------- + +def _cosine_similarity(vec1: list[float], vec2: list[float]) -> float: + """Compute cosine similarity between two vectors.""" + dot = sum(a * b for a, b in zip(vec1, vec2)) + norm1 = math.sqrt(sum(a * a for a in vec1)) + norm2 = math.sqrt(sum(b * b for b in vec2)) + if norm1 == 0 or norm2 == 0: + return 0.0 + return dot / (norm1 * norm2) + + +def _length_score(length: int, opt_min: int, opt_max: int) -> float: + """Score content length on a bell curve centered on [opt_min, opt_max]. + + Returns 1.0 at the optimal midpoint, decaying for shorter/longer content. + """ + if length <= 0: + return 0.0 + midpoint = (opt_min + opt_max) / 2 + # Gaussian-like decay: sigma = (opt_max - opt_min) / 2 + sigma = max((opt_max - opt_min) / 2, 1) + return math.exp(-0.5 * ((length - midpoint) / sigma) ** 2) + + +def _keyword_density(query: str, content: str) -> float: + """Compute what fraction of query terms appear in content. + + Splits query by whitespace, checks each term's presence. + Uses word boundary for Latin text, substring for CJK. + """ + terms = [t.strip() for t in query.split() if t.strip()] + if not terms or not content: + return 0.0 + def _term_in_content(term: str) -> bool: + # CJK characters: substring match (no word boundaries in Chinese/Japanese) + if any('\u4e00' <= ch <= '\u9fff' for ch in term): + return term in content + # Latin text: word boundary match to avoid "art" matching "smart" + return bool(re.search(r'\b' + re.escape(term) + r'\b', content, re.IGNORECASE)) + hits = sum(1 for t in terms if _term_in_content(t)) + return hits / len(terms) + + +def _normalize_minmax(values: list[float]) -> list[float]: + """Min-max normalize a list of values to [0, 1].""" + if not values: + return [] + lo = min(values) + hi = max(values) + if hi - lo < 1e-12: + return [0.5] * len(values) + return [(v - lo) / (hi - lo) for v in values] + + +# --------------------------------------------------------------------------- +# HybridReranker +# --------------------------------------------------------------------------- + +class HybridReranker: + """Multi-signal reranker combining RRF, embedding, intent, and content signals.""" + + def __init__( + self, + config: RerankerConfig | None = None, + embedding_provider: Any | None = None, + ) -> None: + self._config = config or RerankerConfig.default() + self._embedding_provider = embedding_provider + + def rerank( + self, + candidates: list[FusedResult], + query: str, + query_embedding: list[float] | None = None, + intent: str | None = None, + top_k: int = 10, + ) -> list[RerankResult]: + """Rerank candidates using weighted multi-signal fusion. + + Args: + candidates: Fused results from RRF fusion. + query: Original query text (for keyword density). + query_embedding: Query embedding vector (for embedding similarity). + intent: Classified intent string (for intent boost). + top_k: Number of results to return. + + Returns: + Reranked list of RerankResult, sorted by final_score descending. + """ + if not candidates: + return [] + + # Determine which signals are available + has_embedding = ( + query_embedding is not None + and len(query_embedding) > 0 + and self._embedding_provider is not None + ) + # Check if using real (non-fake) embedding + is_real_embedding = has_embedding and _is_real_embedding_provider( + self._embedding_provider + ) + + unavailable: set[str] = set() + if not is_real_embedding: + unavailable.add("embedding_similarity") + + weights = self._config.adjust_weights_for_missing_signals(unavailable) + + # Compute raw signal values for all candidates + rrf_scores = [c.rrf_score for c in candidates] + norm_rrf = _normalize_minmax(rrf_scores) + + # Pre-compute doc embeddings if needed + doc_embeddings: dict[UUID, list[float]] = {} + if is_real_embedding: + doc_embeddings = self._load_doc_embeddings( + [c.chunk_id for c in candidates] + ) + + # Build rerank results + results: list[RerankResult] = [] + for i, cand in enumerate(candidates): + signals: list[RerankSignal] = [] + + # Signal 1: RRF score + w = weights.get("rrf_score", 0.0) + raw = rrf_scores[i] + norm = norm_rrf[i] + signals.append(RerankSignal( + name="rrf_score", weight=w, + raw_value=raw, normalized_value=norm, + contribution=w * norm, + )) + + # Signal 2: Embedding similarity + w = weights.get("embedding_similarity", 0.0) + if is_real_embedding and cand.chunk_id in doc_embeddings: + sim = _cosine_similarity(query_embedding, doc_embeddings[cand.chunk_id]) + else: + sim = 0.0 + signals.append(RerankSignal( + name="embedding_similarity", weight=w, + raw_value=sim, normalized_value=sim, # already in [0,1] + contribution=w * sim, + )) + + # Signal 3: Intent metadata boost + w = weights.get("intent_metadata_boost", 0.0) + boost = self._config.get_intent_boost(intent, cand.doc_type) + # Normalize: boost is already a small positive value, cap at 1.0 + norm_boost = min(boost, 1.0) + signals.append(RerankSignal( + name="intent_metadata_boost", weight=w, + raw_value=boost, normalized_value=norm_boost, + contribution=w * norm_boost, + )) + + # Signal 4: Content quality + w = weights.get("content_quality", 0.0) + cq = self._config.content_quality + len_score = _length_score( + len(cand.content), cq.optimal_length_min, cq.optimal_length_max + ) + kd = _keyword_density(query, cand.content) + content_score = ( + (1 - cq.keyword_density_weight) * len_score + + cq.keyword_density_weight * kd + ) + signals.append(RerankSignal( + name="content_quality", weight=w, + raw_value=content_score, normalized_value=content_score, + contribution=w * content_score, + )) + + # Final score + final = sum(s.contribution for s in signals) + + results.append(RerankResult( + chunk_id=cand.chunk_id, + doc_id=cand.doc_id, + doc_type=cand.doc_type.value if hasattr(cand.doc_type, 'value') else str(cand.doc_type), + content=cand.content, + final_score=final, + signals=signals, + rrf_score=cand.rrf_score, + keyword_rank=cand.keyword_rank, + keyword_contribution=cand.keyword_contribution, + vector_rank=cand.vector_rank, + vector_contribution=cand.vector_contribution, + sources=list(cand.sources), + )) + + # Sort by final_score descending + results.sort(key=lambda r: r.final_score, reverse=True) + + # Assign ranks + for i, r in enumerate(results[:top_k], 1): + r.rank = i + + return results[:top_k] + + def _load_doc_embeddings( + self, chunk_ids: list[UUID] + ) -> dict[UUID, list[float]]: + """Load document embeddings from DB for the given chunk IDs.""" + embeddings: dict[UUID, list[float]] = {} + if not chunk_ids: + return embeddings + try: + from ticketpilot.retrieval.db.connection import get_db_connection # noqa: PLC0415 + + with get_db_connection() as conn: + with conn.cursor() as cur: + placeholders = ",".join(["%s"] * len(chunk_ids)) + cur.execute( + f"SELECT id, embedding FROM knowledge_chunks WHERE id IN ({placeholders})", + [str(cid) for cid in chunk_ids], + ) + for row in cur.fetchall(): + cid = UUID(row[0]) + emb_str = row[1] + if emb_str: + if isinstance(emb_str, str): + emb_str = emb_str.strip("[]") + embeddings[cid] = [float(x) for x in emb_str.split(",")] + elif isinstance(emb_str, list): + embeddings[cid] = [float(x) for x in emb_str] + except Exception as e: + logger.warning("Failed to load document embeddings: %s", e) + return embeddings + + +def _is_real_embedding_provider(provider: Any) -> bool: + """Check if the embedding provider is a real (non-fake) provider.""" + if not hasattr(provider, 'embed') and not hasattr(provider, 'encode'): + return False + name = getattr(provider, "provider_name", "unknown") + return name not in ("fake", "unknown", "") diff --git a/src/ticketpilot/retrieval/pipeline.py b/src/ticketpilot/retrieval/pipeline.py index db4610a..36b8d7c 100644 --- a/src/ticketpilot/retrieval/pipeline.py +++ b/src/ticketpilot/retrieval/pipeline.py @@ -1,17 +1,70 @@ -"""Hybrid retrieval pipeline combining keyword and vector search with RRF fusion.""" +"""Hybrid retrieval pipeline combining keyword and vector search with RRF fusion. +Enhanced with: +- Multi-query expansion (LLM-generated query variants) +- Hybrid reranking (multi-signal weighted fusion) +""" +import logging import time from typing import Optional from ticketpilot.retrieval.keyword_search import keyword_search from ticketpilot.retrieval.providers.fake_embedding import FakeEmbeddingProvider, get_fake_embedding_provider -from ticketpilot.retrieval.reranker import rerank_with_embeddings +from ticketpilot.retrieval.reranker_config import RerankerConfig +from ticketpilot.retrieval.hybrid_reranker import HybridReranker, RerankResult +from ticketpilot.retrieval.query_expander import MultiQueryExpander +from ticketpilot.retrieval.result_merger import merge_retrieval_results + +logger = logging.getLogger(__name__) from ticketpilot.retrieval.rrf import DEFAULT_RRF_K, rrf_fusion from ticketpilot.retrieval.schema.knowledge import DocType -from ticketpilot.retrieval.traces import RetrievalTrace +from ticketpilot.retrieval.traces import FusedResult, RetrievalTrace from ticketpilot.retrieval.vector_search import get_hnsw_params, vector_search +def _single_query_retrieval( + query: str, + query_embedding: list[float], + top_k: int, + doc_types: Optional[list[DocType]], + exclude_business_domains: Optional[list[str]], + embedding_provider, + rrf_k: int, +) -> tuple[list[FusedResult], list, str, int, list, int]: + """Run keyword + vector + RRF for a single query. + + Returns (fused_results, keyword_results, search_method, keyword_latency, + vector_results, vector_latency). + """ + provider_name = getattr(embedding_provider, "provider_name", "unknown") + + # Keyword search + kw_start = time.perf_counter() + kw_results, kw_method = keyword_search( + query=query, + top_k=top_k * 2, + doc_types=doc_types, + exclude_business_domains=exclude_business_domains, + ) + kw_latency = int((time.perf_counter() - kw_start) * 1000) + + # Vector search + vec_start = time.perf_counter() + vec_results, _ = vector_search( + query_embedding=query_embedding, + top_k=top_k * 2, + doc_types=doc_types, + exclude_business_domains=exclude_business_domains, + embedding_provider_name=provider_name, + ) + vec_latency = int((time.perf_counter() - vec_start) * 1000) + + # RRF fusion + fused = rrf_fusion(keyword_results=kw_results, vector_results=vec_results, k=rrf_k) + + return fused, kw_results, kw_method, kw_latency, vec_results, vec_latency + + def hybrid_retrieval( query: str, top_k: int = 10, @@ -19,18 +72,21 @@ def hybrid_retrieval( exclude_business_domains: Optional[list[str]] = None, embedding_provider: Optional[FakeEmbeddingProvider] = None, rrf_k: int = DEFAULT_RRF_K, - enable_reranking: bool = True, # Enabled with improved strategy + enable_reranking: bool = True, + # New params for hybrid reranking (backward compatible) + intent: Optional[str] = None, + enable_query_expansion: bool = False, + reranker_config: Optional[RerankerConfig] = None, ) -> RetrievalTrace: """ Perform hybrid retrieval combining keyword and vector search. Pipeline: - 1. Generate query embedding using the embedding provider - 2. Run keyword search (FTS + LIKE fallback) - 3. Run vector search (HNSW) - 4. Fuse results using RRF - 5. Re-rank top results using embedding similarity (optional) - 6. Return complete trace for debugging and audit + 1. (Optional) Expand query into variants via LLM + 2. For each query variant: keyword search + vector search + RRF fusion + 3. Merge results from all variants (sum_score dedup) + 4. (Optional) Hybrid rerank with multi-signal fusion + 5. Return complete trace for debugging and audit Args: query: Search query string @@ -39,98 +95,175 @@ def hybrid_retrieval( embedding_provider: Embedding provider (default: FakeEmbeddingProvider) rrf_k: RRF k parameter (default: 60) enable_reranking: Enable re-ranking step (default: True) + intent: Classified intent string (for intent-aware reranking) + enable_query_expansion: Enable LLM-based query expansion (default: False) + reranker_config: Custom reranker config (default: from YAML or built-in) Returns: RetrievalTrace with complete pipeline information """ total_start_time = time.perf_counter() - # Use provided embedding provider or default (matching DB dimension) + # Use provided embedding provider or default if embedding_provider is None: - from ticketpilot.retrieval.vector_search import _detect_embedding_dim + from ticketpilot.retrieval.vector_search import _detect_embedding_dim # noqa: PLC0415 dim = _detect_embedding_dim() embedding_provider = get_fake_embedding_provider(dimension=dim) - # Generate query embedding + provider_name = getattr(embedding_provider, "provider_name", "unknown") + is_real = provider_name not in ("fake", "unknown", "") + + # Generate query embedding for the original query query_embedding = embedding_provider.embed(query) - # Keyword search - keyword_start = time.perf_counter() - keyword_results, keyword_search_method = keyword_search( - query=query, - top_k=top_k * 2, # Fetch more to account for fusion - doc_types=doc_types, - exclude_business_domains=exclude_business_domains, - ) - keyword_latency_ms = int((time.perf_counter() - keyword_start) * 1000) + # --- Step 0: Query Expansion --- + expansion_start = time.perf_counter() + query_variants = [query] + expansion_latency = 0 + if enable_query_expansion: + expander = MultiQueryExpander() + query_variants = expander.expand(query, intent or "") + expansion_latency = int((time.perf_counter() - expansion_start) * 1000) - # Get provider name for trace (handles both FakeEmbeddingProvider and others) - provider_name = getattr(embedding_provider, "provider_name", "unknown") + # --- Step 1-3: Per-query retrieval + RRF --- + all_fused: list[list[FusedResult]] = [] + # Use the first query's keyword/vector results for the trace + first_kw_results = [] + first_kw_method = "fts" + first_kw_latency = 0 + first_vec_results = [] + first_vec_latency = 0 - # Vector search - vector_start = time.perf_counter() - vector_results, vector_latency_ms = vector_search( - query_embedding=query_embedding, - top_k=top_k * 2, # Fetch more to account for fusion - doc_types=doc_types, - exclude_business_domains=exclude_business_domains, - embedding_provider_name=provider_name, - ) - vector_latency_ms = int((time.perf_counter() - vector_start) * 1000) - - # RRF Fusion - fusion_start = time.perf_counter() - fused_results = rrf_fusion( - keyword_results=keyword_results, - vector_results=vector_results, - k=rrf_k, - ) - fusion_latency_ms = int((time.perf_counter() - fusion_start) * 1000) + for i, q in enumerate(query_variants): + # Generate embedding for variant (reuse original for first query) + if i == 0: + q_emb = query_embedding + else: + q_emb = embedding_provider.embed(q) - # Re-ranking (optional) - rerank_latency_ms = 0 - if enable_reranking and fused_results: + fused, kw_res, kw_meth, kw_lat, vec_res, vec_lat = _single_query_retrieval( + query=q, + query_embedding=q_emb, + top_k=top_k, + doc_types=doc_types, + exclude_business_domains=exclude_business_domains, + embedding_provider=embedding_provider, + rrf_k=rrf_k, + ) + all_fused.append(fused) + + if i == 0: + first_kw_results = kw_res + first_kw_method = kw_meth + first_kw_latency = kw_lat + first_vec_results = vec_res + first_vec_latency = vec_lat + + # --- Step 4: Merge results from all query variants --- + merge_start = time.perf_counter() + if len(all_fused) > 1: + merged = merge_retrieval_results(all_fused, strategy="sum_score") + else: + merged = all_fused[0] if all_fused else [] + merge_latency = int((time.perf_counter() - merge_start) * 1000) + merged_count = len(merged) + + # --- Step 5: Reranking --- + rerank_latency = 0 + rerank_signals_data = None + reranker_weights_data = None + final_fused: list[FusedResult] = [] + + if enable_reranking and merged: rerank_start = time.perf_counter() - - # Take top 20 for re-ranking (more than final top_k) - candidates = fused_results[:20] - - # Re-rank using embedding similarity - fused_results = rerank_with_embeddings( + + # Load reranker config + if reranker_config is None: + try: + reranker_config = RerankerConfig.from_yaml("config/reranker.yaml") + except Exception as e: + logger.warning("Failed to load reranker config from YAML, using default: %s", e) + reranker_config = RerankerConfig.default() + + # Take top candidates for reranking + candidates = merged[: max(top_k * 3, 20)] + + reranker = HybridReranker( + config=reranker_config, + embedding_provider=embedding_provider, + ) + reranked: list[RerankResult] = reranker.rerank( + candidates=candidates, + query=query, query_embedding=query_embedding, - fused_results=candidates, + intent=intent, top_k=top_k, - embedding_provider=embedding_provider, ) - rerank_latency_ms = int((time.perf_counter() - rerank_start) * 1000) + + # Convert RerankResult back to FusedResult for downstream compatibility + final_fused = [r.to_fused_result() for r in reranked] + + # Extract trace data + if reranked: + rerank_signals_data = [] + for r in reranked: + sig_data = { + "chunk_id": str(r.chunk_id), + "final_score": round(r.final_score, 6), + "signals": [ + { + "name": s.name, + "weight": round(s.weight, 4), + "raw": round(s.raw_value, 6), + "normalized": round(s.normalized_value, 6), + "contribution": round(s.contribution, 6), + } + for s in r.signals + ], + } + rerank_signals_data.append(sig_data) + + # Get actual weights from the first result's signals + if reranked[0].signals: + reranker_weights_data = { + s.name: round(s.weight, 4) for s in reranked[0].signals + } + + rerank_latency = int((time.perf_counter() - rerank_start) * 1000) else: - # Limit to top_k without re-ranking - fused_results = fused_results[:top_k] + final_fused = merged[:top_k] - final_evidence_ids = [r.chunk_id for r in fused_results] + final_evidence_ids = [r.chunk_id for r in final_fused] # Total latency - total_latency_ms = int((time.perf_counter() - total_start_time) * 1000) + total_latency = int((time.perf_counter() - total_start_time) * 1000) # Build trace trace = RetrievalTrace( query=query, query_embedding=query_embedding, - keyword_results=keyword_results, - keyword_latency_ms=keyword_latency_ms, - keyword_search_method=keyword_search_method, - vector_results=vector_results, - vector_latency_ms=vector_latency_ms, - fused_results=fused_results, - fusion_latency_ms=fusion_latency_ms, + keyword_results=first_kw_results, + keyword_latency_ms=first_kw_latency, + keyword_search_method=first_kw_method, + vector_results=first_vec_results, + vector_latency_ms=first_vec_latency, + fused_results=final_fused, + fusion_latency_ms=merge_latency, rrf_k=rrf_k, final_evidence_ids=final_evidence_ids, - total_latency_ms=total_latency_ms, + total_latency_ms=total_latency, embedding_provider=provider_name, hnsw_params=get_hnsw_params(), top_k=top_k, - rerank_latency_ms=rerank_latency_ms, + rerank_latency_ms=rerank_latency, reranking_enabled=enable_reranking, + # Hybrid reranking fields + query_variants=query_variants if len(query_variants) > 1 else None, + expansion_latency_ms=expansion_latency, + merged_result_count=merged_count, + rerank_signals=rerank_signals_data, + reranker_weights=reranker_weights_data, + has_real_embedding=is_real, ) return trace @@ -145,14 +278,6 @@ def simple_retrieval( Simple retrieval interface returning just content. Convenience function for cases where trace is not needed. - - Args: - query: Search query string - top_k: Maximum number of results - doc_types: Optional filter by document types - - Returns: - List of content strings for top-k results """ trace = hybrid_retrieval(query, top_k, doc_types) - return [r.content for r in trace.fused_results] \ No newline at end of file + return [r.content for r in trace.fused_results] diff --git a/src/ticketpilot/retrieval/query_expander.py b/src/ticketpilot/retrieval/query_expander.py new file mode 100644 index 0000000..b3cc156 --- /dev/null +++ b/src/ticketpilot/retrieval/query_expander.py @@ -0,0 +1,136 @@ +"""Multi-query expansion using LLM for improved retrieval recall. + +Generates query variants by asking an LLM to rephrase the original query +from different semantic angles. Falls back to original query on failure. +""" +from __future__ import annotations + +import json +import logging +import os +import re + +logger = logging.getLogger(__name__) + +_EXPANSION_PROMPT = """\ +你是一个搜索查询优化器。给定一个客服工单查询,生成 {n} 个不同角度的搜索关键词变体。 +要求: +- 每个变体 5-15 个字 +- 覆盖不同语义角度(同义词、上位词、具体场景) +- 不要重复原始查询 +- 只输出 JSON 数组,不要解释 + +查询:{query} +意图:{intent} + +输出格式:["变体1", "变体2"] +""" + + +class MultiQueryExpander: + """Generate query variants using LLM for improved recall. + + Uses the same LLM endpoint as DraftAgent (configured via env vars). + Falls back to returning only the original query on any failure. + """ + + def __init__( + self, + num_variants: int = 2, + base_url: str | None = None, + api_key: str | None = None, + model: str | None = None, + timeout: int = 15, + ) -> None: + self._num_variants = num_variants + self._base_url = ( + base_url + or os.environ.get("TICKETPILOT_LLM_BASE_URL", "https://api.deepseek.com") + ).rstrip("/") + self._api_key = api_key or os.environ.get("TICKETPILOT_LLM_API_KEY", "") + self._model = model or os.environ.get("TICKETPILOT_LLM_MODEL", "deepseek-chat") + self._timeout = timeout + + def expand(self, query: str, intent: str = "") -> list[str]: + """Return [original_query, variant_1, variant_2, ...]. + + On any failure, returns [original_query] only. + """ + if not self._api_key: + logger.debug("No LLM API key, skipping query expansion") + return [query] + + try: + variants = self._call_llm(query, intent) + # Validate variants + valid = [v for v in variants if self._is_valid_variant(v, query)] + result = [query] + valid[: self._num_variants] + logger.info( + "Query expansion: original_len=%d -> %d variants (total valid: %d)", + len(query), len(valid[: self._num_variants]), len(valid), + ) + return result + except Exception as e: + logger.warning("Query expansion failed, using original: %s", e) + return [query] + + def _call_llm(self, query: str, intent: str) -> list[str]: + """Call LLM to generate query variants.""" + import urllib.request # noqa: PLC0415 + + prompt = _EXPANSION_PROMPT.format( + n=self._num_variants, query=query, intent=intent + ) + payload = { + "model": self._model, + "messages": [{"role": "user", "content": prompt}], + "max_tokens": 200, + "temperature": 0.5, + } + req = urllib.request.Request( + f"{self._base_url}/chat/completions", + data=json.dumps(payload).encode("utf-8"), + headers={ + "Authorization": f"Bearer {self._api_key}", + "Content-Type": "application/json", + }, + method="POST", + ) + with urllib.request.urlopen(req, timeout=self._timeout) as resp: + result = json.loads(resp.read().decode("utf-8")) + + content = result.get("choices", [{}])[0].get("message", {}).get("content", "") + return self._parse_variants(content) + + def _parse_variants(self, text: str) -> list[str]: + """Extract JSON array from LLM response.""" + # Try markdown code fence + match = re.search(r"```(?:json)?\s*\n?(.*?)\n?```", text, re.DOTALL) + if match: + try: + parsed = json.loads(match.group(1).strip()) + if isinstance(parsed, list): + return [str(v) for v in parsed] + except json.JSONDecodeError: + pass + + # Try raw JSON array + match = re.search(r"\[.*\]", text, re.DOTALL) + if match: + try: + parsed = json.loads(match.group(0)) + if isinstance(parsed, list): + return [str(v) for v in parsed] + except json.JSONDecodeError: + pass + + return [] + + def _is_valid_variant(self, variant: str, original: str) -> bool: + """Check if a variant is valid: non-empty, different from original, reasonable length.""" + v = variant.strip() + if not v or len(v) > 50: + return False + if v == original.strip(): + return False + return True diff --git a/src/ticketpilot/retrieval/reranker_config.py b/src/ticketpilot/retrieval/reranker_config.py new file mode 100644 index 0000000..b22c7f8 --- /dev/null +++ b/src/ticketpilot/retrieval/reranker_config.py @@ -0,0 +1,164 @@ +"""Configuration for the hybrid reranker. + +Defines signal weights, intent boost tables, and content quality parameters. +Supports loading from YAML files for A/B experiments. +""" +from __future__ import annotations + +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any + + +@dataclass +class ContentQualityConfig: + """Parameters for the content quality scoring signal.""" + optimal_length_min: int = 200 + optimal_length_max: int = 800 + keyword_density_weight: float = 0.5 + + def __post_init__(self) -> None: + if self.optimal_length_min > self.optimal_length_max: + raise ValueError( + f"optimal_length_min ({self.optimal_length_min}) must be <= " + f"optimal_length_max ({self.optimal_length_max})" + ) + if not 0 <= self.keyword_density_weight <= 1: + raise ValueError( + f"keyword_density_weight must be between 0 and 1, " + f"got {self.keyword_density_weight}" + ) + + +@dataclass +class RerankerConfig: + """Configuration for the hybrid reranker. + + Attributes: + weights: Signal name -> weight mapping. Must sum to 1.0. + intent_boost_table: IntentClass value -> {doc_type: boost_value}. + content_quality: Content quality scoring parameters. + num_query_variants: Number of LLM-generated query variants. + """ + + weights: dict[str, float] = field(default_factory=dict) + intent_boost_table: dict[str, dict[str, float]] = field(default_factory=dict) + content_quality: ContentQualityConfig = field(default_factory=ContentQualityConfig) + num_query_variants: int = 2 + + # --- Validation --- + + def validate(self) -> None: + """Validate config. Raises ValueError on issues.""" + if not self.weights: + raise ValueError("weights cannot be empty") + total = sum(self.weights.values()) + if abs(total - 1.0) > 1e-6: + raise ValueError( + f"weights must sum to 1.0, got {total:.4f}: {self.weights}" + ) + for name, w in self.weights.items(): + if w < 0: + raise ValueError(f"weight '{name}' must be >= 0, got {w}") + + # --- Factories --- + + @classmethod + def default(cls) -> RerankerConfig: + """Default config with balanced weights.""" + cfg = cls( + weights={ + "rrf_score": 0.40, + "embedding_similarity": 0.25, + "intent_metadata_boost": 0.20, + "content_quality": 0.15, + }, + intent_boost_table={ + "refund": {"policy": 0.15, "faq": 0.10}, + "return_exchange": {"policy": 0.15, "faq": 0.10}, + "account_issue": {"policy": 0.10, "faq": 0.10}, + "technical_issue": {"faq": 0.10, "case": 0.10}, + "product_consulting": {"faq": 0.15}, + "logistics": {"faq": 0.10, "case": 0.10}, + "complaint": {"case": 0.15, "policy": 0.10}, + "other": {}, + }, + content_quality=ContentQualityConfig( + optimal_length_min=200, + optimal_length_max=800, + keyword_density_weight=0.5, + ), + num_query_variants=2, + ) + cfg.validate() + return cfg + + @classmethod + def from_yaml(cls, path: str | Path) -> RerankerConfig: + """Load config from a YAML file. + + Falls back to default if file not found. + """ + import logging # noqa: PLC0415 + import yaml # noqa: PLC0415 + + logger = logging.getLogger(__name__) + path = Path(path) + if not path.exists(): + logger.warning("Reranker config file not found: %s, using defaults", path) + return cls.default() + + with path.open("r", encoding="utf-8") as f: + data: dict[str, Any] = yaml.safe_load(f) or {} + + weights = data.get("weights", {}) + intent_boost = data.get("intent_boost", {}) + cq_data = data.get("content_quality", {}) + num_variants = data.get("num_query_variants", 2) + + cq = ContentQualityConfig( + optimal_length_min=cq_data.get("optimal_length_min", 200), + optimal_length_max=cq_data.get("optimal_length_max", 800), + keyword_density_weight=cq_data.get("keyword_density_weight", 0.5), + ) + + cfg = cls( + weights=weights, + intent_boost_table=intent_boost, + content_quality=cq, + num_query_variants=num_variants, + ) + cfg.validate() + return cfg + + # --- Helpers --- + + def get_intent_boost(self, intent: str | None, doc_type: str) -> float: + """Get boost value for an intent+doc_type combination.""" + if intent is None: + return 0.0 + # Normalize doc_type to lowercase for table lookup + dt_lower = doc_type.lower() if doc_type else "" + return self.intent_boost_table.get(intent, {}).get(dt_lower, 0.0) + + def adjust_weights_for_missing_signals( + self, unavailable_signals: set[str] + ) -> dict[str, float]: + """Redistribute weights when signals are unavailable. + + Returns a new weight dict with unavailable signals removed + and remaining weights renormalized to sum to 1.0. + """ + if not unavailable_signals: + return dict(self.weights) + + available = { + k: v for k, v in self.weights.items() if k not in unavailable_signals + } + total = sum(available.values()) + if total <= 0: + # Fallback: equal weight on all available + n = len(available) or 1 + return {k: 1.0 / n for k in available} + + return {k: v / total for k, v in available.items()} diff --git a/src/ticketpilot/retrieval/result_merger.py b/src/ticketpilot/retrieval/result_merger.py new file mode 100644 index 0000000..3cc422f --- /dev/null +++ b/src/ticketpilot/retrieval/result_merger.py @@ -0,0 +1,151 @@ +"""Merge retrieval results from multiple query variants. + +Supports three merge strategies: +- max_score: keep highest RRF score per chunk_id +- sum_score: sum RRF scores across queries (boosts docs found by multiple queries) +- rrf_again: treat each query as a ranker, apply second-level RRF +""" +from __future__ import annotations + +import logging +from collections import defaultdict +from uuid import UUID + +from ticketpilot.retrieval.traces import FusedResult + +logger = logging.getLogger(__name__) + + +def merge_retrieval_results( + result_sets: list[list[FusedResult]], + strategy: str = "sum_score", +) -> list[FusedResult]: + """Merge multiple retrieval result sets, deduplicating by chunk_id. + + Args: + result_sets: List of FusedResult lists, one per query variant. + strategy: Merge strategy - "max_score", "sum_score", or "rrf_again". + + Returns: + Merged and deduplicated list of FusedResult, sorted by score descending. + """ + if not result_sets: + return [] + + # Flatten and filter empty sets + non_empty = [rs for rs in result_sets if rs] + if not non_empty: + return [] + if len(non_empty) == 1: + return list(non_empty[0]) + + if strategy == "sum_score": + return _merge_sum_score(non_empty) + elif strategy == "max_score": + return _merge_max_score(non_empty) + elif strategy == "rrf_again": + return _merge_rrf_again(non_empty) + else: + logger.warning("Unknown merge strategy '%s', falling back to 'sum_score'", strategy) + return _merge_sum_score(non_empty) + + +def _merge_sum_score( + result_sets: list[list[FusedResult]], +) -> list[FusedResult]: + """Sum RRF scores for the same chunk_id across query variants. + + Docs found by multiple queries get higher scores (multi-path validation). + """ + best: dict[UUID, FusedResult] = {} + score_sums: dict[UUID, float] = defaultdict(float) + + for result_set in result_sets: + for r in result_set: + score_sums[r.chunk_id] += r.rrf_score + # Keep the version with most info (prefer one with both keyword+vector) + if r.chunk_id not in best or r.rrf_score > best[r.chunk_id].rrf_score: + best[r.chunk_id] = r + + # Build merged results with summed scores + merged: list[FusedResult] = [] + for cid, representative in best.items(): + merged.append(FusedResult( + chunk_id=cid, + doc_id=representative.doc_id, + doc_type=representative.doc_type, + content=representative.content, + rrf_score=score_sums[cid], + keyword_rank=representative.keyword_rank, + keyword_contribution=representative.keyword_contribution, + vector_rank=representative.vector_rank, + vector_contribution=representative.vector_contribution, + sources=representative.sources + (["multi_query"] if "multi_query" not in representative.sources else []), + )) + + merged.sort(key=lambda r: r.rrf_score, reverse=True) + return merged + + +def _merge_max_score( + result_sets: list[list[FusedResult]], +) -> list[FusedResult]: + """Keep the highest RRF score per chunk_id.""" + best: dict[UUID, FusedResult] = {} + + for result_set in result_sets: + for r in result_set: + if r.chunk_id not in best or r.rrf_score > best[r.chunk_id].rrf_score: + best[r.chunk_id] = r + + merged = list(best.values()) + merged.sort(key=lambda r: r.rrf_score, reverse=True) + return merged + + +def _merge_rrf_again( + result_sets: list[list[FusedResult]], +) -> list[FusedResult]: + """Apply second-level RRF: treat each query variant as a ranker. + + Uses RRF k=60 on the rank positions within each query's results. + """ + from ticketpilot.retrieval.rrf import DEFAULT_RRF_K # noqa: PLC0415 + k = DEFAULT_RRF_K + # Build per-query rank maps + rank_maps: list[dict[UUID, int]] = [] + representative: dict[UUID, FusedResult] = {} + + for result_set in result_sets: + rank_map: dict[UUID, int] = {} + for i, r in enumerate(result_set, 1): + rank_map[r.chunk_id] = i + if r.chunk_id not in representative: + representative[r.chunk_id] = r + rank_maps.append(rank_map) + + # Compute second-level RRF scores + rrf_scores: dict[UUID, float] = defaultdict(float) + for rm in rank_maps: + for cid, rank in rm.items(): + rrf_scores[cid] += 1.0 / (k + rank) + + # Build merged results + merged: list[FusedResult] = [] + for cid, score in rrf_scores.items(): + rep = representative[cid] + merged.append(FusedResult( + chunk_id=cid, + doc_id=rep.doc_id, + doc_type=rep.doc_type, + content=rep.content, + rrf_score=score, + keyword_rank=rep.keyword_rank, + keyword_contribution=rep.keyword_contribution, + vector_rank=rep.vector_rank, + vector_contribution=rep.vector_contribution, + sources=rep.sources + ["rrf_again"], + )) + + merged.sort(key=lambda r: r.rrf_score, reverse=True) + return merged diff --git a/src/ticketpilot/retrieval/retrieve_evidence.py b/src/ticketpilot/retrieval/retrieve_evidence.py index f193c67..5692777 100644 --- a/src/ticketpilot/retrieval/retrieve_evidence.py +++ b/src/ticketpilot/retrieval/retrieve_evidence.py @@ -6,6 +6,7 @@ from ticketpilot.retrieval.pipeline import hybrid_retrieval from ticketpilot.retrieval.providers.fake_embedding import FakeEmbeddingProvider from ticketpilot.retrieval.query_builder import build_retrieval_query +from ticketpilot.retrieval.reranker_config import RerankerConfig from ticketpilot.retrieval.schema.knowledge import DocType from ticketpilot.retrieval.traces import RetrievalTrace from ticketpilot.schema.evidence import EvidenceCandidate @@ -19,14 +20,26 @@ def retrieve_evidence( top_k: int = 10, doc_types: list[DocType] | None = None, embedding_provider: Optional[FakeEmbeddingProvider] = None, + # New params for hybrid reranking (backward compatible) + enable_query_expansion: bool = False, + reranker_config: Optional[RerankerConfig] = None, ) -> tuple[list[EvidenceCandidate], RetrievalTrace]: """Retrieve evidence candidates from the knowledge base. Constructs a retrieval query from ticket state, runs hybrid - retrieval, and maps fused results to evidence candidates. + retrieval with optional query expansion and hybrid reranking, + and maps fused results to evidence candidates. Always returns a RetrievalTrace, even when no results are found. """ query = build_retrieval_query(normalized_text, intent, risk_flags) - trace = hybrid_retrieval(query=query, top_k=top_k, doc_types=doc_types, embedding_provider=embedding_provider) + trace = hybrid_retrieval( + query=query, + top_k=top_k, + doc_types=doc_types, + embedding_provider=embedding_provider, + intent=intent.value if intent else None, + enable_query_expansion=enable_query_expansion, + reranker_config=reranker_config, + ) candidates = map_fused_to_evidence(trace.fused_results) return candidates, trace diff --git a/src/ticketpilot/retrieval/traces.py b/src/ticketpilot/retrieval/traces.py index af25f29..ff17bd2 100644 --- a/src/ticketpilot/retrieval/traces.py +++ b/src/ticketpilot/retrieval/traces.py @@ -1,6 +1,6 @@ """Retrieval trace schema for debugging and explainability.""" -from datetime import datetime, timezone, timezone +from datetime import datetime, timezone from typing import Any, Optional from uuid import UUID @@ -191,6 +191,34 @@ class RetrievalTrace(BaseModel): description="Whether re-ranking was enabled", ) + # Hybrid reranking metadata + query_variants: Optional[list[str]] = Field( + default=None, + description="Query variants used for multi-query expansion", + ) + expansion_latency_ms: int = Field( + default=0, + ge=0, + description="Query expansion latency in milliseconds", + ) + merged_result_count: int = Field( + default=0, + ge=0, + description="Number of results after merge+dedup, before reranking", + ) + rerank_signals: Optional[list[dict[str, Any]]] = Field( + default=None, + description="Per-result signal breakdown from hybrid reranker", + ) + reranker_weights: Optional[dict[str, float]] = Field( + default=None, + description="Actual weights used by hybrid reranker (may be adjusted)", + ) + has_real_embedding: bool = Field( + default=False, + description="Whether a real (non-fake) embedding provider was used", + ) + def get_result_by_chunk_id(self, chunk_id: UUID) -> Optional[FusedResult]: """Get fused result by chunk ID.""" for result in self.fused_results: diff --git a/tests/unit/test_hybrid_reranker.py b/tests/unit/test_hybrid_reranker.py new file mode 100644 index 0000000..5e95323 --- /dev/null +++ b/tests/unit/test_hybrid_reranker.py @@ -0,0 +1,246 @@ +"""Unit tests for HybridReranker.""" +import math +from unittest.mock import MagicMock +from uuid import uuid4 + +import pytest + +from ticketpilot.retrieval.hybrid_reranker import ( + HybridReranker, + RerankResult, + _cosine_similarity, + _keyword_density, + _length_score, + _normalize_minmax, +) +from ticketpilot.retrieval.reranker_config import RerankerConfig +from ticketpilot.retrieval.schema.knowledge import DocType +from ticketpilot.retrieval.traces import FusedResult + + +def _make_fused( + chunk_id=None, doc_type="FAQ", content="test content", rrf_score=0.5 +) -> FusedResult: + return FusedResult( + chunk_id=chunk_id or uuid4(), + doc_id=uuid4(), + doc_type=DocType(doc_type), + content=content, + rrf_score=rrf_score, + keyword_rank=1, + keyword_contribution=0.016, + vector_rank=2, + vector_contribution=0.015, + sources=["keyword", "vector"], + ) + + +class TestCosineSimilarity: + def test_identical_vectors(self): + v = [1.0, 0.0, 0.0] + assert abs(_cosine_similarity(v, v) - 1.0) < 1e-9 + + def test_orthogonal_vectors(self): + assert abs(_cosine_similarity([1, 0], [0, 1])) < 1e-9 + + def test_opposite_vectors(self): + assert abs(_cosine_similarity([1, 0], [-1, 0]) - (-1.0)) < 1e-9 + + def test_zero_vector(self): + assert _cosine_similarity([0, 0], [1, 0]) == 0.0 + + +class TestLengthScore: + def test_optimal_length_scores_highest(self): + score = _length_score(500, 200, 800) + assert score > 0.9 + + def test_very_short_scores_low(self): + score = _length_score(10, 200, 800) + assert score < 0.3 + + def test_very_long_scores_lower(self): + score_optimal = _length_score(500, 200, 800) + score_long = _length_score(3000, 200, 800) + assert score_optimal > score_long + + def test_zero_length(self): + assert _length_score(0, 200, 800) == 0.0 + + +class TestKeywordDensity: + def test_all_terms_present(self): + assert _keyword_density("退款 到账", "退款一直没到账怎么办") == 1.0 + + def test_partial_terms(self): + assert _keyword_density("退款 到账 物流", "退款政策说明") == pytest.approx(1 / 3) + + def test_no_terms(self): + assert _keyword_density("退款", "物流发货说明") == 0.0 + + def test_empty_query(self): + assert _keyword_density("", "some content") == 0.0 + + +class TestNormalizeMinMax: + def test_uniform_values(self): + result = _normalize_minmax([5, 5, 5]) + assert result == [0.5, 0.5, 0.5] + + def test_normal_range(self): + result = _normalize_minmax([0, 5, 10]) + assert result[0] == 0.0 + assert result[1] == pytest.approx(0.5) + assert result[2] == 1.0 + + def test_empty(self): + assert _normalize_minmax([]) == [] + + +class TestHybridReranker: + def test_empty_candidates(self): + reranker = HybridReranker() + assert reranker.rerank([], "test") == [] + + def test_single_candidate(self): + c = _make_fused(rrf_score=0.5, content="退款政策说明 退款条件") + reranker = HybridReranker() + results = reranker.rerank([c], "退款", intent="refund", top_k=5) + assert len(results) == 1 + assert results[0].rank == 1 + assert results[0].final_score > 0 + + def test_ranking_order(self): + c1 = _make_fused(rrf_score=0.3, content="退款政策说明") + c2 = _make_fused(rrf_score=0.8, content="退款退款退款退款退款") + reranker = HybridReranker() + results = reranker.rerank([c1, c2], "退款", top_k=10) + assert len(results) == 2 + # c2 has higher RRF + more keyword hits, should rank first + assert results[0].chunk_id == c2.chunk_id + assert results[0].rank == 1 + + def test_intent_boost_effect(self): + policy_doc = _make_fused( + doc_type="POLICY", rrf_score=0.3, content="退款政策 退款条件" + ) + case_doc = _make_fused( + doc_type="CASE", rrf_score=0.3, content="退款案例 退款处理" + ) + reranker = HybridReranker() + results = reranker.rerank( + [policy_doc, case_doc], "退款", intent="refund", top_k=10 + ) + # Policy doc should rank higher due to intent boost for refund→policy + policy_result = next(r for r in results if r.chunk_id == policy_doc.chunk_id) + case_result = next(r for r in results if r.chunk_id == case_doc.chunk_id) + assert policy_result.final_score > case_result.final_score + + def test_signals_recorded(self): + c = _make_fused(content="退款政策说明 退款流程") + reranker = HybridReranker() + results = reranker.rerank([c], "退款", top_k=5) + assert len(results[0].signals) == 4 + signal_names = {s.name for s in results[0].signals} + assert signal_names == { + "rrf_score", "embedding_similarity", + "intent_metadata_boost", "content_quality", + } + + def test_fake_embedding_weight_redistribution(self): + """With fake embedding provider, embedding weight should be 0.""" + cfg = RerankerConfig.default() + # No embedding_provider = fake + reranker = HybridReranker(config=cfg, embedding_provider=None) + c = _make_fused(content="test content") + results = reranker.rerank([c], "test", top_k=5) + # embedding_similarity signal should have weight=0 + emb_signal = next( + s for s in results[0].signals if s.name == "embedding_similarity" + ) + assert emb_signal.weight == 0.0 + + def test_to_fused_result_conversion(self): + c = _make_fused(content="退款政策") + reranker = HybridReranker() + results = reranker.rerank([c], "退款", top_k=5) + fused = results[0].to_fused_result() + assert isinstance(fused, FusedResult) + assert "hybrid_rerank" in fused.sources + assert fused.chunk_id == c.chunk_id + + def test_top_k_less_than_candidates(self): + """top_k truncates results.""" + candidates = [_make_fused(rrf_score=0.1 * i) for i in range(5)] + reranker = HybridReranker() + results = reranker.rerank(candidates, "test", top_k=2) + assert len(results) == 2 + assert results[0].rank == 1 + assert results[1].rank == 2 + + def test_top_k_more_than_candidates(self): + """top_k > len(candidates) returns all candidates.""" + candidates = [_make_fused(rrf_score=0.5)] + reranker = HybridReranker() + results = reranker.rerank(candidates, "test", top_k=10) + assert len(results) == 1 + + +class TestIsRealEmbeddingProvider: + def test_real_provider(self): + from ticketpilot.retrieval.hybrid_reranker import _is_real_embedding_provider + provider = MagicMock() + provider.embed = MagicMock() + provider.provider_name = "openai" + assert _is_real_embedding_provider(provider) is True + + def test_fake_provider(self): + from ticketpilot.retrieval.hybrid_reranker import _is_real_embedding_provider + provider = MagicMock() + provider.embed = MagicMock() + provider.provider_name = "fake" + assert _is_real_embedding_provider(provider) is False + + def test_no_embed_or_encode(self): + from ticketpilot.retrieval.hybrid_reranker import _is_real_embedding_provider + provider = MagicMock(spec=[]) # no attributes + assert _is_real_embedding_provider(provider) is False + + def test_unknown_provider_name(self): + from ticketpilot.retrieval.hybrid_reranker import _is_real_embedding_provider + provider = MagicMock() + provider.encode = MagicMock() + # No provider_name attribute → getattr returns "unknown" + del provider.provider_name + assert _is_real_embedding_provider(provider) is False + + def test_none_provider(self): + from ticketpilot.retrieval.hybrid_reranker import _is_real_embedding_provider + assert _is_real_embedding_provider(None) is False + + def test_encode_method_sufficient(self): + from ticketpilot.retrieval.hybrid_reranker import _is_real_embedding_provider + provider = MagicMock(spec=["encode", "provider_name"]) + provider.encode = MagicMock() + provider.provider_name = "bge" + assert _is_real_embedding_provider(provider) is True + + +class TestKeywordDensityEdgeCases: + def test_latin_word_boundary_no_false_positive(self): + """'art' should NOT match inside 'smart'.""" + assert _keyword_density("art", "smart car") == 0.0 + + def test_latin_word_boundary_exact_match(self): + """'art' matches standalone 'Art'.""" + assert _keyword_density("art", "Art of war") == 1.0 + + def test_cjk_substring_match(self): + """CJK terms use substring matching.""" + assert _keyword_density("退款", "退款政策说明") == 1.0 + + def test_mixed_cjk_latin(self): + """Mixed query: CJK substring + Latin word boundary.""" + assert _keyword_density("退款 policy", "退款 policy 说明") == 1.0 + # 'policy' inside 'policyholder' should not match + assert _keyword_density("policy", "policyholder agreement") == 0.0 diff --git a/tests/unit/test_query_expander.py b/tests/unit/test_query_expander.py new file mode 100644 index 0000000..4deadcf --- /dev/null +++ b/tests/unit/test_query_expander.py @@ -0,0 +1,95 @@ +"""Unit tests for MultiQueryExpander.""" +import json +from unittest.mock import MagicMock, patch + +import pytest + +from ticketpilot.retrieval.query_expander import MultiQueryExpander + + +class TestMultiQueryExpander: + def test_no_api_key_returns_original(self, monkeypatch): + monkeypatch.delenv("TICKETPILOT_LLM_API_KEY", raising=False) + expander = MultiQueryExpander(api_key="") + result = expander.expand("退款没到账", "refund") + assert result == ["退款没到账"] + + def test_parse_json_array(self): + expander = MultiQueryExpander(api_key="fake") + variants = expander._parse_variants('["退款进度", "退款到账时间"]') + assert variants == ["退款进度", "退款到账时间"] + + def test_parse_markdown_fence(self): + expander = MultiQueryExpander(api_key="fake") + text = '```json\n["退款进度", "退款到账时间"]\n```' + variants = expander._parse_variants(text) + assert variants == ["退款进度", "退款到账时间"] + + def test_parse_invalid_returns_empty(self): + expander = MultiQueryExpander(api_key="fake") + assert expander._parse_variants("not json") == [] + + def test_is_valid_variant_normal(self): + expander = MultiQueryExpander(api_key="fake") + assert expander._is_valid_variant("退款进度查询", "退款没到账") + + def test_is_valid_variant_empty(self): + expander = MultiQueryExpander(api_key="fake") + assert not expander._is_valid_variant("", "query") + + def test_is_valid_variant_same_as_original(self): + expander = MultiQueryExpander(api_key="fake") + assert not expander._is_valid_variant("退款没到账", "退款没到账") + + def test_is_valid_variant_too_long(self): + expander = MultiQueryExpander(api_key="fake") + assert not expander._is_valid_variant("a" * 51, "query") + + @patch("ticketpilot.retrieval.query_expander.MultiQueryExpander._call_llm") + def test_expand_success(self, mock_llm): + mock_llm.return_value = ["退款进度", "退款到账时间"] + expander = MultiQueryExpander(api_key="fake-key") + result = expander.expand("退款没到账", "refund") + assert result == ["退款没到账", "退款进度", "退款到账时间"] + + @patch("ticketpilot.retrieval.query_expander.MultiQueryExpander._call_llm") + def test_expand_llm_failure_fallback(self, mock_llm): + mock_llm.side_effect = RuntimeError("API error") + expander = MultiQueryExpander(api_key="fake-key") + result = expander.expand("退款没到账", "refund") + assert result == ["退款没到账"] + + @patch("ticketpilot.retrieval.query_expander.MultiQueryExpander._call_llm") + def test_expand_filters_invalid_variants(self, mock_llm): + mock_llm.return_value = ["", "a" * 51, "有效变体"] + expander = MultiQueryExpander(api_key="fake-key") + result = expander.expand("退款没到账") + assert result == ["退款没到账", "有效变体"] + + @patch("ticketpilot.retrieval.query_expander.MultiQueryExpander._call_llm") + def test_expand_default_intent(self, mock_llm): + """Test that expand works with default intent parameter (empty string).""" + mock_llm.return_value = ["退款进度"] + expander = MultiQueryExpander(api_key="fake-key") + result = expander.expand("退款没到账") # no intent arg + assert result == ["退款没到账", "退款进度"] + + @patch("ticketpilot.retrieval.query_expander.MultiQueryExpander._call_llm") + def test_expand_respects_num_variants_limit(self, mock_llm): + """When LLM returns more variants than num_variants, truncate.""" + mock_llm.return_value = ["变体1", "变体2", "变体3"] + expander = MultiQueryExpander(api_key="fake-key", num_variants=2) + result = expander.expand("original query") + assert len(result) == 3 # original + 2 variants + assert "变体3" not in result + + def test_is_valid_variant_with_whitespace(self): + """Variant that equals original after strip should be invalid.""" + expander = MultiQueryExpander(api_key="fake") + assert not expander._is_valid_variant(" 退款没到账 ", "退款没到账") + + def test_parse_non_string_elements(self): + """JSON array with non-string elements should convert to strings.""" + expander = MultiQueryExpander(api_key="fake") + variants = expander._parse_variants('[123, true, "正常"]') + assert variants == ["123", "True", "正常"] diff --git a/tests/unit/test_reranker_config.py b/tests/unit/test_reranker_config.py new file mode 100644 index 0000000..27357aa --- /dev/null +++ b/tests/unit/test_reranker_config.py @@ -0,0 +1,143 @@ +"""Unit tests for RerankerConfig.""" +import tempfile +from pathlib import Path + +import pytest + +from ticketpilot.retrieval.reranker_config import ContentQualityConfig, RerankerConfig + + +class TestRerankerConfigValidation: + def test_default_config_is_valid(self): + cfg = RerankerConfig.default() + cfg.validate() # should not raise + + def test_weights_sum_to_one(self): + cfg = RerankerConfig.default() + total = sum(cfg.weights.values()) + assert abs(total - 1.0) < 1e-6 + + def test_empty_weights_raises(self): + cfg = RerankerConfig(weights={}) + with pytest.raises(ValueError, match="weights cannot be empty"): + cfg.validate() + + def test_weights_not_summing_raises(self): + cfg = RerankerConfig(weights={"a": 0.3, "b": 0.3}) + with pytest.raises(ValueError, match="must sum to 1.0"): + cfg.validate() + + def test_negative_weight_raises(self): + cfg = RerankerConfig(weights={"a": 1.5, "b": -0.5}) + with pytest.raises(ValueError, match="must be >= 0"): + cfg.validate() + + +class TestRerankerConfigFromYaml: + def test_load_from_yaml(self, tmp_path): + yaml_content = """\ +weights: + rrf_score: 0.50 + embedding_similarity: 0.30 + intent_metadata_boost: 0.10 + content_quality: 0.10 + +intent_boost: + refund: + policy: 0.20 + +content_quality: + optimal_length_min: 100 + optimal_length_max: 600 + keyword_density_weight: 0.3 + +num_query_variants: 3 +""" + p = tmp_path / "reranker.yaml" + p.write_text(yaml_content) + cfg = RerankerConfig.from_yaml(str(p)) + assert cfg.weights["rrf_score"] == 0.50 + assert cfg.num_query_variants == 3 + assert cfg.content_quality.optimal_length_min == 100 + + def test_missing_file_returns_default(self): + cfg = RerankerConfig.from_yaml("/nonexistent/path.yaml") + assert cfg.weights == RerankerConfig.default().weights + + +class TestRerankerConfigHelpers: + def test_get_intent_boost_match(self): + cfg = RerankerConfig.default() + boost = cfg.get_intent_boost("refund", "policy") + assert boost == 0.15 + + def test_get_intent_boost_no_match(self): + cfg = RerankerConfig.default() + boost = cfg.get_intent_boost("refund", "case") + assert boost == 0.0 + + def test_get_intent_boost_none_intent(self): + cfg = RerankerConfig.default() + boost = cfg.get_intent_boost(None, "policy") + assert boost == 0.0 + + def test_adjust_weights_no_missing(self): + cfg = RerankerConfig.default() + adjusted = cfg.adjust_weights_for_missing_signals(set()) + assert adjusted == cfg.weights + + def test_adjust_weights_embedding_missing(self): + cfg = RerankerConfig.default() + adjusted = cfg.adjust_weights_for_missing_signals({"embedding_similarity"}) + assert "embedding_similarity" not in adjusted + total = sum(adjusted.values()) + assert abs(total - 1.0) < 1e-6 + + def test_adjust_weights_preserves_proportions(self): + cfg = RerankerConfig.default() + adjusted = cfg.adjust_weights_for_missing_signals({"embedding_similarity"}) + # rrf was 0.40, now should be 0.40/0.75 ≈ 0.533 + expected_ratio = 0.40 / 0.75 + assert abs(adjusted["rrf_score"] - expected_ratio) < 1e-6 + + def test_adjust_weights_all_signals_missing(self): + """When all signals are removed, return empty dict.""" + cfg = RerankerConfig.default() + all_signals = set(cfg.weights.keys()) + adjusted = cfg.adjust_weights_for_missing_signals(all_signals) + assert isinstance(adjusted, dict) + assert len(adjusted) == 0 + + +class TestContentQualityConfig: + def test_valid_config(self): + cq = ContentQualityConfig(optimal_length_min=100, optimal_length_max=500) + assert cq.optimal_length_min == 100 + + def test_min_greater_than_max_raises(self): + with pytest.raises(ValueError, match="optimal_length_min.*must be <="): + ContentQualityConfig(optimal_length_min=800, optimal_length_max=200) + + def test_density_weight_out_of_range_raises(self): + with pytest.raises(ValueError, match="keyword_density_weight must be between 0 and 1"): + ContentQualityConfig(keyword_density_weight=1.5) + + def test_negative_density_weight_raises(self): + with pytest.raises(ValueError, match="keyword_density_weight must be between 0 and 1"): + ContentQualityConfig(keyword_density_weight=-0.1) + + +class TestRerankerConfigFromYamlEdgeCases: + def test_malformed_yaml_raises(self, tmp_path): + """Invalid YAML syntax should raise.""" + p = tmp_path / "bad.yaml" + p.write_text(":\n invalid: [yaml\n") + with pytest.raises(Exception): + RerankerConfig.from_yaml(str(p)) + + def test_empty_yaml_raises_validation(self, tmp_path): + """Empty YAML file produces empty weights → validation error.""" + p = tmp_path / "empty.yaml" + p.write_text("") + with pytest.raises(ValueError, match="weights cannot be empty"): + RerankerConfig.from_yaml(str(p)) diff --git a/tests/unit/test_result_merger.py b/tests/unit/test_result_merger.py new file mode 100644 index 0000000..de6a7d9 --- /dev/null +++ b/tests/unit/test_result_merger.py @@ -0,0 +1,130 @@ +"""Unit tests for result_merger.""" +from uuid import uuid4 + +import pytest + +from ticketpilot.retrieval.result_merger import merge_retrieval_results +from ticketpilot.retrieval.schema.knowledge import DocType +from ticketpilot.retrieval.traces import FusedResult + + +def _fused(chunk_id=None, rrf_score=0.5, content="test", sources=None): + return FusedResult( + chunk_id=chunk_id or uuid4(), + doc_id=uuid4(), + doc_type=DocType.FAQ, + content=content, + rrf_score=rrf_score, + keyword_rank=1, + keyword_contribution=0.016, + sources=sources or ["keyword"], + ) + + +class TestMergeRetrievalResults: + def test_empty_input(self): + assert merge_retrieval_results([]) == [] + + def test_all_empty_sets(self): + assert merge_retrieval_results([[], []]) == [] + + def test_single_set_passthrough(self): + r = _fused() + result = merge_retrieval_results([[r]]) + assert len(result) == 1 + assert result[0].chunk_id == r.chunk_id + + def test_sum_score_dedup(self): + cid = uuid4() + r1 = _fused(chunk_id=cid, rrf_score=0.3, sources=["keyword"]) + r2 = _fused(chunk_id=cid, rrf_score=0.2, sources=["vector"]) + merged = merge_retrieval_results([[r1], [r2]], strategy="sum_score") + assert len(merged) == 1 + assert merged[0].rrf_score == pytest.approx(0.5) + + def test_sum_score_different_chunks(self): + c1 = uuid4() + c2 = uuid4() + r1 = _fused(chunk_id=c1, rrf_score=0.3) + r2 = _fused(chunk_id=c2, rrf_score=0.5) + merged = merge_retrieval_results([[r1], [r2]], strategy="sum_score") + assert len(merged) == 2 + # c2 should rank first (higher score) + assert merged[0].chunk_id == c2 + + def test_max_score_strategy(self): + cid = uuid4() + r1 = _fused(chunk_id=cid, rrf_score=0.3) + r2 = _fused(chunk_id=cid, rrf_score=0.7) + merged = merge_retrieval_results([[r1], [r2]], strategy="max_score") + assert len(merged) == 1 + assert merged[0].rrf_score == pytest.approx(0.7) + + def test_rrf_again_strategy(self): + c1 = uuid4() + c2 = uuid4() + # c1 ranked #1 in both queries, c2 ranked #2 + r1q1 = _fused(chunk_id=c1, rrf_score=0.5) + r1q2 = _fused(chunk_id=c1, rrf_score=0.4) + r2q1 = _fused(chunk_id=c2, rrf_score=0.3) + r2q2 = _fused(chunk_id=c2, rrf_score=0.6) + merged = merge_retrieval_results( + [[r1q1, r2q1], [r1q2, r2q2]], strategy="rrf_again" + ) + assert len(merged) == 2 + # c1 ranked higher in both, should be first + assert merged[0].chunk_id == c1 + + def test_multi_query_marker(self): + r = _fused() + merged = merge_retrieval_results([[r], [r]]) + assert "multi_query" in merged[0].sources + + def test_unknown_strategy_defaults_to_sum_score(self): + """Unknown strategy string falls back to sum_score.""" + cid = uuid4() + r1 = _fused(chunk_id=cid, rrf_score=0.3) + r2 = _fused(chunk_id=cid, rrf_score=0.2) + merged = merge_retrieval_results([[r1], [r2]], strategy="unknown_strategy") + assert len(merged) == 1 + assert merged[0].rrf_score == pytest.approx(0.5) # sum_score behavior + + def test_sum_score_prefers_highest_rrf_representative(self): + """When same chunk appears multiple times, representative has highest rrf_score.""" + cid = uuid4() + r1 = _fused(chunk_id=cid, rrf_score=0.1, sources=["keyword"]) + r2 = _fused(chunk_id=cid, rrf_score=0.8, sources=["vector"]) + merged = merge_retrieval_results([[r1], [r2]], strategy="sum_score") + assert len(merged) == 1 + # Representative should be r2 (higher rrf_score) + assert "vector" in merged[0].sources + + def test_multi_query_marker_no_duplicate(self): + """Same chunk from 3 queries should have only one 'multi_query' marker.""" + cid = uuid4() + r1 = _fused(chunk_id=cid, rrf_score=0.3, sources=["keyword"]) + r2 = _fused(chunk_id=cid, rrf_score=0.2, sources=["keyword"]) + r3 = _fused(chunk_id=cid, rrf_score=0.1, sources=["keyword"]) + merged = merge_retrieval_results([[r1], [r2], [r3]], strategy="sum_score") + assert len(merged) == 1 + assert merged[0].sources.count("multi_query") == 1 + + def test_rrf_again_precise_scores(self): + """Verify exact RRF scores with k=60.""" + c1 = uuid4() + c2 = uuid4() + r1q1 = _fused(chunk_id=c1, rrf_score=0.5) + r1q2 = _fused(chunk_id=c1, rrf_score=0.4) + r2q1 = _fused(chunk_id=c2, rrf_score=0.3) + r2q2 = _fused(chunk_id=c2, rrf_score=0.6) + merged = merge_retrieval_results( + [[r1q1, r2q1], [r1q2, r2q2]], strategy="rrf_again" + ) + k = 60 + expected_c1 = 2 * (1 / (k + 1)) # Both rank 1 + expected_c2 = 2 * (1 / (k + 2)) # Both rank 2 + assert len(merged) == 2 + assert merged[0].chunk_id == c1 + assert merged[0].rrf_score == pytest.approx(expected_c1) + assert merged[1].chunk_id == c2 + assert merged[1].rrf_score == pytest.approx(expected_c2)