diff --git a/.github/workflows/build-embedding-image.yml b/.github/workflows/build-embedding-image.yml new file mode 100644 index 0000000..d00bf8f --- /dev/null +++ b/.github/workflows/build-embedding-image.yml @@ -0,0 +1,54 @@ +name: Build embedding image + +on: + push: + branches: [ main, Phins-branch ] + pull_request: + branches: [ main ] + +jobs: + build: + runs-on: ubuntu-latest + env: + IMAGE_NAME: ghcr.io/${{ github.repository_owner }}/lexcam-embedding-service + MODEL_NAME: intfloat/multilingual-e5-small + MODEL_URL: https://github.com/${{ github.repository }}/releases/download/${{ vars.EMBEDDING_MODEL_RELEASE_TAG }}/model.onnx + MODEL_SHA256: ${{ vars.EMBEDDING_MODEL_SHA256 }} + + steps: + - name: Checkout + uses: actions/checkout@v4 + + - name: Set up Python + uses: actions/setup-python@v4 + with: + python-version: '3.11' + + - name: Build Docker image + run: | + docker build \ + --tag $IMAGE_NAME:${{ github.sha }} \ + --build-arg MODEL_NAME=$MODEL_NAME \ + --build-arg MODEL_URL=$MODEL_URL \ + --build-arg MODEL_SHA256=$MODEL_SHA256 \ + services/embedding-service + + - name: Smoke test image + run: | + id=$(docker create --name tmp_test -p 8001:8000 $IMAGE_NAME:${{ github.sha }}) + docker start $id + sleep 5 + docker ps --filter "name=tmp_test" --format "{{.Names}} {{.Status}}" + docker exec tmp_test /bin/sh -c "curl -sS -f http://localhost:8000/api/v1/health || exit 1" + docker stop $id + docker rm $id + + - name: Publish to GHCR + if: github.event_name == 'push' && secrets.GHCR_PAT + env: + CR_PAT: ${{ secrets.GHCR_PAT }} + run: | + echo $CR_PAT | docker login ghcr.io -u ${{ github.repository_owner }} --password-stdin + docker tag $IMAGE_NAME:${{ github.sha }} $IMAGE_NAME:latest + docker push $IMAGE_NAME:${{ github.sha }} + docker push $IMAGE_NAME:latest diff --git a/.gitignore b/.gitignore index 129478d..4140414 100644 --- a/.gitignore +++ b/.gitignore @@ -53,6 +53,9 @@ coverage.xml # Docker *.tar +# Downloaded or generated embedding artifacts +services/embedding-service/model/*.onnx + # Helm charts/*.tgz diff --git a/README.md b/README.md index c398468..6199700 100644 --- a/README.md +++ b/README.md @@ -1 +1,15 @@ # LexCamAI + +## Embedding model storage + +The embedding service now expects the ONNX model to live outside git, with a +GitHub Releases asset as the preferred source. + +Use a release tag such as `embedding-model-v1`, upload the model as +`model.onnx`, and set these repository variables for the build workflow: + +- `EMBEDDING_MODEL_RELEASE_TAG` +- `EMBEDDING_MODEL_SHA256` + +The image build downloads the asset from: +`https://github.com///releases/download//model.onnx` diff --git a/docker-compose.dev.yml b/docker-compose.dev.yml index ade93ae..6a30cf3 100644 --- a/docker-compose.dev.yml +++ b/docker-compose.dev.yml @@ -62,6 +62,8 @@ services: embedding-service: build: context: ./services/embedding-service + args: + - MODEL_NAME=intfloat/multilingual-e5-small container_name: lexcam-embedding-service environment: SERVICE_NAME: embedding-service @@ -80,6 +82,28 @@ services: volumes: - embedding_cache:/cache + knowledge-base-service: + build: + context: ./services/knowledge-base-service + container_name: lexcam-knowledge-base-service + environment: + SERVICE_NAME: knowledge-base-service + API_PREFIX: /api/v1 + DATABASE_URL: postgresql+psycopg://lexcam:lexcam_dev@postgres:5432/lexcam_knowledge + QDRANT_URL: http://qdrant:6333 + EMBEDDING_SERVICE_URL: http://embedding-service:8000 + REDIS_URL: redis://redis:6379/0 + LOG_LEVEL: INFO + PLAIN_SUMMARY_CACHE_TTL_SECONDS: 604800 + MAX_SEARCH_RESULTS: 10 + ports: + - "8003:8000" + depends_on: + - postgres + - qdrant + - redis + - embedding-service + volumes: postgres_data: qdrant_data: diff --git a/phins-changes.patch b/phins-changes.patch new file mode 100644 index 0000000..a73b35e Binary files /dev/null and b/phins-changes.patch differ diff --git a/scripts/init-databases.sql b/scripts/init-databases.sql index 36281d9..50ad80d 100644 --- a/scripts/init-databases.sql +++ b/scripts/init-databases.sql @@ -1,7 +1,7 @@ -- LexCam PostgreSQL initialization -- Runs automatically on first container start. -- Design choice: database-per-service on a single Postgres instance --- (separate logical databases like lexcam_users, lexcam_rag_sessions, etc.). +-- (separate logical databases like lexcam_users, lexcam_rag_sessions). CREATE DATABASE lexcam_users; CREATE DATABASE lexcam_lawyers; diff --git a/scripts/populate_embedding_cache.ps1 b/scripts/populate_embedding_cache.ps1 new file mode 100644 index 0000000..47dc5e3 --- /dev/null +++ b/scripts/populate_embedding_cache.ps1 @@ -0,0 +1,23 @@ +Param( + [string]$Image = 'lexcam/embedding-service:from-test' +) + +Write-Output "Using image: $Image" + +$tmp = New-TemporaryFile +$tmpDir = Split-Path $tmp -Parent + +try { + $cid = (docker create $Image).Trim() + Write-Output "Created container $cid" + docker cp "$cid`:/cache/model.onnx" "$tmpDir\model.onnx" + docker rm $cid | Out-Null + + # create temporary container with volume mounted + $vcontainer = (docker create --name tmp_embedding_volume -v embedding_cache:/cache busybox).Trim() + docker cp "$tmpDir\model.onnx" "$vcontainer`:/cache/model.onnx" + docker rm $vcontainer | Out-Null + Write-Output "Model copied into volume 'embedding_cache'" +} finally { + Remove-Item "$tmpDir\model.onnx" -ErrorAction SilentlyContinue +} diff --git a/scripts/populate_embedding_cache.sh b/scripts/populate_embedding_cache.sh new file mode 100644 index 0000000..636292c --- /dev/null +++ b/scripts/populate_embedding_cache.sh @@ -0,0 +1,32 @@ +#!/usr/bin/env bash +set -euo pipefail + +# Populate the named Docker volume `embedding_cache` with model files +# Usage: ./scripts/populate_embedding_cache.sh [image] +# Default image: lexcam/embedding-service:from-test + +IMAGE=${1:-lexcam/embedding-service:from-test} + +echo "Using image: $IMAGE" + +TMPDIR=$(mktemp -d) +cleanup() { + rm -rf "$TMPDIR" +} +trap cleanup EXIT + +echo "Creating temporary container from image to extract model..." +CID=$(docker create "$IMAGE") +echo "Copying /cache/model.onnx from container $CID to host temp" +docker cp "$CID":/cache/model.onnx "$TMPDIR"/model.onnx +docker rm "$CID" >/dev/null + +echo "Creating temporary container with embedding_cache volume mounted..." +VCONTAINER=$(docker create --name tmp_embedding_volume -v embedding_cache:/cache busybox) +echo "Copying model into volume" +docker cp "$TMPDIR"/model.onnx tmp_embedding_volume:/cache/model.onnx +docker rm tmp_embedding_volume >/dev/null + +echo "Model copied into volume 'embedding_cache'" + +echo "Done" diff --git a/services/embedding-service/Dockerfile b/services/embedding-service/Dockerfile index 840ca37..90600a9 100644 --- a/services/embedding-service/Dockerfile +++ b/services/embedding-service/Dockerfile @@ -1,10 +1,15 @@ -FROM python:3.11-slim +FROM mcr.microsoft.com/devcontainers/python:1-3.11-bookworm ENV PYTHONDONTWRITEBYTECODE=1 \ PYTHONUNBUFFERED=1 \ TRANSFORMERS_CACHE=/cache \ HF_HOME=/cache +ARG MODEL_NAME=intfloat/multilingual-e5-small +ENV MODEL_NAME=${MODEL_NAME} +ARG MODEL_URL +ARG MODEL_SHA256= + WORKDIR /app RUN useradd --create-home appuser \ @@ -14,6 +19,41 @@ RUN useradd --create-home appuser \ COPY requirements.txt . RUN pip install --no-cache-dir -r requirements.txt +RUN python3 - <<'PY' +import hashlib +import os +import pathlib +import urllib.request + +model_url = os.environ.get("MODEL_URL", "").strip() +if not model_url: + raise SystemExit( + "MODEL_URL build arg is required. Point it at the externally stored model artifact." + ) + +expected_sha256 = os.environ.get("MODEL_SHA256", "").strip().lower() +target = pathlib.Path("/cache/model.onnx") +target.parent.mkdir(parents=True, exist_ok=True) + +hasher = hashlib.sha256() if expected_sha256 else None +with urllib.request.urlopen(model_url) as response, target.open("wb") as output_file: + while True: + chunk = response.read(1024 * 1024) + if not chunk: + break + output_file.write(chunk) + if hasher: + hasher.update(chunk) + +if hasher and hasher.hexdigest().lower() != expected_sha256: + target.unlink(missing_ok=True) + raise SystemExit( + f"MODEL_SHA256 mismatch for {target}: expected {expected_sha256}, got {hasher.hexdigest().lower()}" + ) + +print(f"Downloaded external model artifact to {target}") +PY + COPY app ./app EXPOSE 8000 diff --git a/services/embedding-service/app/api/v1/routes.py b/services/embedding-service/app/api/v1/routes.py index 7350073..2d08b0d 100644 --- a/services/embedding-service/app/api/v1/routes.py +++ b/services/embedding-service/app/api/v1/routes.py @@ -1,7 +1,5 @@ -from __future__ import annotations - import time -from fastapi import APIRouter, Depends, HTTPException, Request +from fastapi import APIRouter, Body, Depends, HTTPException, Request from app.config import settings from app.limiting import rate_limit @@ -29,7 +27,7 @@ async def health(request: Request) -> HealthResponse: dependencies=[Depends(verify_api_key)], ) @rate_limit() -async def embed(request: Request, payload: EmbeddingRequest) -> EmbeddingResponse: +async def embed(request: Request, payload: EmbeddingRequest = Body(...)) -> EmbeddingResponse: if len(payload.texts) > settings.max_batch_size: raise HTTPException( status_code=413, diff --git a/services/embedding-service/app/schemas.py b/services/embedding-service/app/schemas.py index 5430fb8..0a986cb 100644 --- a/services/embedding-service/app/schemas.py +++ b/services/embedding-service/app/schemas.py @@ -2,11 +2,11 @@ from typing import Literal -from pydantic import BaseModel, Field, conlist +from pydantic import BaseModel, Field class EmbeddingRequest(BaseModel): - texts: conlist(str, min_items=1) = Field(..., description="Texts to embed") + texts: list[str] = Field(..., min_items=1, description="Texts to embed") input_type: Literal["query", "passage"] | None = Field( default=None, description="Optional E5 prefix to apply", diff --git a/services/embedding-service/app/services/embedding.py b/services/embedding-service/app/services/embedding.py index 8de8252..6c7f27e 100644 --- a/services/embedding-service/app/services/embedding.py +++ b/services/embedding-service/app/services/embedding.py @@ -1,23 +1,30 @@ from __future__ import annotations import logging +from typing import Optional import anyio -from sentence_transformers import SentenceTransformer +import numpy as np +import onnxruntime as ort +from transformers import AutoTokenizer from app.config import settings +from pathlib import Path class EmbeddingModel: - def __init__(self) -> None: - self._model: SentenceTransformer | None = None - self._dimension: int | None = None + """ONNX-based embedding model runner. - @property - def model(self) -> SentenceTransformer: - if self._model is None: - raise RuntimeError("Embedding model is not loaded") - return self._model + Expects an ONNX model file at `/cache/model.onnx`. This class uses a + Hugging Face tokenizer (from `settings.model_name`) and runs the ONNX + session to obtain the last hidden state, applies mean-pooling and + optional L2-normalization to produce sentence embeddings. + """ + + def __init__(self) -> None: + self._tokenizer: Optional[AutoTokenizer] = None + self._session: Optional[ort.InferenceSession] = None + self._dimension: Optional[int] = None @property def dimension(self) -> int: @@ -26,42 +33,72 @@ def dimension(self) -> int: return self._dimension async def load(self) -> None: - def _load() -> SentenceTransformer: - model = SentenceTransformer(settings.model_name, device=settings.device) - model.max_seq_length = settings.max_seq_length - return model + def _load(): + tokenizer = AutoTokenizer.from_pretrained(settings.model_name, use_fast=True) + + model_path = "/cache/model.onnx" + if not Path(model_path).exists(): + raise RuntimeError( + "ONNX model not found at /cache/model.onnx.\n" + "Provide the model as a release asset and let the Docker build download it via MODEL_URL." + ) + session = ort.InferenceSession(model_path, providers=["CPUExecutionProvider"]) + + # Attempt to infer hidden dimension from the first output shape + outputs = session.get_outputs() + hidden_size: Optional[int] = None + if outputs: + out_shape = outputs[0].shape # e.g., (batch, seq_len, hidden) + if len(out_shape) >= 3 and out_shape[2] is not None: + hidden_size = int(out_shape[2]) + + return tokenizer, session, hidden_size - self._model = await anyio.to_thread.run_sync(_load) - self._dimension = self._model.get_sentence_embedding_dimension() + self._tokenizer, self._session, self._dimension = await anyio.to_thread.run_sync(_load) logging.getLogger(__name__).info( - "Embedding model loaded: %s (%s, dim=%s)", - settings.model_name, - settings.device, - self._dimension, + "ONNX embedding model loaded: %s (dim=%s)", settings.model_name, self._dimension ) - async def embed( - self, - texts: list[str], - input_type: str | None, - normalize: bool | None, - ) -> list[list[float]]: + async def embed(self, texts: list[str], input_type: str | None, normalize: bool | None) -> list[list[float]]: if input_type: prefix = f"{input_type}: " - texts = [f"{prefix}{text}" for text in texts] + texts = [f"{prefix}{t}" for t in texts] - normalize_embeddings = ( - settings.normalize_embeddings if normalize is None else normalize - ) + normalize_embeddings = settings.normalize_embeddings if normalize is None else normalize def _encode() -> list[list[float]]: - vectors = self.model.encode( + assert self._tokenizer is not None and self._session is not None + + enc = self._tokenizer( texts, - batch_size=min(len(texts), settings.max_batch_size), - convert_to_numpy=True, - normalize_embeddings=normalize_embeddings, - show_progress_bar=False, + padding=True, + truncation=True, + max_length=settings.max_seq_length, + return_tensors="np", ) - return vectors.tolist() + + # ONNX runtime expects numpy inputs + ort_inputs = {k: v for k, v in enc.items()} + outputs = self._session.run(None, ort_inputs) + + # assume outputs[0] is last_hidden_state: (batch, seq_len, hidden) + last_hidden = outputs[0] + attention_mask = enc.get("attention_mask") + if attention_mask is None: + # fallback: average across tokens + embeddings = last_hidden.mean(axis=1) + else: + mask = attention_mask.astype(np.float32)[..., None] + summed = (last_hidden * mask).sum(axis=1) + counts = mask.sum(axis=1) + counts[counts == 0] = 1 + embeddings = summed / counts + + if normalize_embeddings: + norms = np.linalg.norm(embeddings, axis=1, keepdims=True) + norms[norms == 0] = 1 + embeddings = embeddings / norms + + return embeddings.tolist() return await anyio.to_thread.run_sync(_encode) diff --git a/services/embedding-service/convert.Dockerfile b/services/embedding-service/convert.Dockerfile new file mode 100644 index 0000000..1308529 --- /dev/null +++ b/services/embedding-service/convert.Dockerfile @@ -0,0 +1,16 @@ +FROM python:3.11-slim + +ENV PYTHONDONTWRITEBYTECODE=1 \ + PYTHONUNBUFFERED=1 + +WORKDIR /workspace + +RUN apt-get update \ + && apt-get install -y --no-install-recommends build-essential git curl ca-certificates \ + && rm -rf /var/lib/apt/lists/* \ + && pip install --no-cache-dir --upgrade pip \ + && pip install --no-cache-dir torch --index-url https://download.pytorch.org/whl/cpu transformers onnx onnxruntime optimum sentence-transformers + +COPY services/embedding-service/scripts/convert_to_onnx.py /workspace/convert_to_onnx.py + +ENTRYPOINT ["python", "/workspace/convert_to_onnx.py"] diff --git a/services/embedding-service/model/.gitkeep b/services/embedding-service/model/.gitkeep new file mode 100644 index 0000000..0519ecb --- /dev/null +++ b/services/embedding-service/model/.gitkeep @@ -0,0 +1 @@ + \ No newline at end of file diff --git a/services/embedding-service/requirements.txt b/services/embedding-service/requirements.txt index 4888dcc..b72289a 100644 --- a/services/embedding-service/requirements.txt +++ b/services/embedding-service/requirements.txt @@ -2,8 +2,8 @@ fastapi==0.111.0 uvicorn[standard]==0.30.1 gunicorn==22.0.0 pydantic-settings==2.3.4 -sentence-transformers==3.0.1 -torch==2.3.1 --index-url https://download.pytorch.org/whl/cpu +transformers==4.35.0 numpy==1.26.4 anyio==4.4.0 slowapi==0.1.9 +onnxruntime==1.26.0 diff --git a/services/embedding-service/scripts/convert_to_onnx.py b/services/embedding-service/scripts/convert_to_onnx.py new file mode 100644 index 0000000..e97c77d --- /dev/null +++ b/services/embedding-service/scripts/convert_to_onnx.py @@ -0,0 +1,98 @@ +"""Helper to convert a Hugging Face / Sentence-Transformers model to ONNX. + +This script attempts to use the transformers ONNX conversion tooling. Converting +requires PyTorch and the conversion tools; run this locally and upload the +generated `model.onnx` file as an external artifact, such as a GitHub Releases +asset, before building the Docker image. + +Usage (example / recommended): + +1. Create a Python venv with PyTorch and transformers installed: + + pip install torch sentence-transformers transformers onnx onnxruntime + +2. Run this script: + + python services/embedding-service/scripts/convert_to_onnx.py --model intfloat/multilingual-e5-small --output model.onnx + +Notes: +- Conversion commands and compatibility depend on the model architecture. +- If this script cannot perform an automated conversion on your platform, + follow Hugging Face or sentence-transformers conversion guides to produce + an ONNX file and upload it as your release asset. +""" + +import argparse +import shutil +import subprocess +import sys +from pathlib import Path + + +def export_with_torch(model_id: str, output: Path) -> None: + import torch + from transformers import AutoModel, AutoTokenizer + + tokenizer = AutoTokenizer.from_pretrained(model_id, use_fast=True) + model = AutoModel.from_pretrained(model_id) + model.eval() + + sample = tokenizer( + ["This is a sample text for ONNX export."], + padding=True, + truncation=True, + max_length=128, + return_tensors="pt", + ) + + output.parent.mkdir(parents=True, exist_ok=True) + if output.exists(): + output.unlink() + data_file = Path(str(output) + ".data") + if data_file.exists(): + data_file.unlink() + + torch.onnx.export( + model, + (sample["input_ids"], sample["attention_mask"]), + str(output), + input_names=["input_ids", "attention_mask"], + output_names=["last_hidden_state"], + dynamic_axes={ + "input_ids": {0: "batch_size", 1: "sequence"}, + "attention_mask": {0: "batch_size", 1: "sequence"}, + "last_hidden_state": {0: "batch_size", 1: "sequence"}, + }, + opset_version=18, + do_constant_folding=True, + external_data=False, + ) + if data_file.exists(): + data_file.unlink() + print("Torch ONNX export complete; wrote:", output) + + +def main(): + p = argparse.ArgumentParser() + p.add_argument("--model", required=True, help="Hugging Face model id") + p.add_argument("--output", required=True, help="Path to write model.onnx") + args = p.parse_args() + + output = Path(args.output) + output.parent.mkdir(parents=True, exist_ok=True) + + # Try using the transformers.onnx CLI if available + cli = shutil.which("transformers-onnx") or shutil.which("transformers.onnx") + if cli: + print("Found transformers ONNX CLI; attempting conversion using it...") + cmd = [sys.executable, "-m", "transformers.onnx", "--model", args.model, str(output)] + subprocess.check_call(cmd) + print("Conversion complete; wrote:", output) + return + + print("transformers ONNX CLI not found; falling back to torch.onnx export.") + export_with_torch(args.model, output) + + +if __name__ == "__main__": + main() diff --git a/services/knowledge-base-service/Dockerfile b/services/knowledge-base-service/Dockerfile new file mode 100644 index 0000000..5e58202 --- /dev/null +++ b/services/knowledge-base-service/Dockerfile @@ -0,0 +1,22 @@ +FROM python:3.11-slim + +ENV PYTHONDONTWRITEBYTECODE=1 \ + PYTHONUNBUFFERED=1 + +WORKDIR /app + +RUN useradd --create-home appuser + +COPY requirements.txt . +RUN pip install --no-cache-dir -r requirements.txt + +COPY app ./app +COPY scripts ./scripts + +EXPOSE 8000 + +USER appuser + +ENV WEB_CONCURRENCY=1 + +CMD ["gunicorn", "-k", "uvicorn.workers.UvicornWorker", "app.main:app", "--bind", "0.0.0.0:8000", "--workers", "1", "--timeout", "120"] diff --git a/services/knowledge-base-service/README.md b/services/knowledge-base-service/README.md new file mode 100644 index 0000000..19ef60e --- /dev/null +++ b/services/knowledge-base-service/README.md @@ -0,0 +1,40 @@ +Knowledge Base Service + +Quick start (development with docker-compose): + +1. Ensure docker-compose.dev.yml includes the `knowledge-base-service` entry (it does by default). +2. Start dependencies: Postgres, Qdrant, Redis, Embedding Service via docker-compose: + +```bash +docker compose -f docker-compose.dev.yml up -d postgres qdrant redis embedding-service +``` + +3. Create the KB schema in the `lexcam_knowledge` database (run inside the Postgres container or via psql): + +```bash +# from repo root +cat services/knowledge-base-service/sql/init_kb_schema.sql | docker exec -i lexcam-postgres psql -U lexcam -d lexcam_knowledge +``` + +4. Build and start the Knowledge Base Service: + +```bash +docker compose -f docker-compose.dev.yml up -d knowledge-base-service +``` + +5. Seed a sample document/article into the KB (make sure services are running): + +```bash +# from repo root +python services/knowledge-base-service/scripts/seed_kb.py +``` + +6. Test the public search endpoint: + +```bash +curl -X POST http://localhost:8003/api/v1/search -H 'Content-Type: application/json' -d '{"query":"santé travailleurs"}' +``` + +Notes: +- The service uses SQLAlchemy `create_all` at startup, but you still need to run the provided SQL to create the tsvector index and trigger for full-text search. +- The Indexing Worker and ingestion workflows will use the `/internal/retrieve` endpoint and Qdrant upserts. diff --git a/services/knowledge-base-service/app/__init__.py b/services/knowledge-base-service/app/__init__.py new file mode 100644 index 0000000..0b16126 --- /dev/null +++ b/services/knowledge-base-service/app/__init__.py @@ -0,0 +1 @@ +# LexCam Knowledge Base Service package diff --git a/services/knowledge-base-service/app/api/v1/__init__.py b/services/knowledge-base-service/app/api/v1/__init__.py new file mode 100644 index 0000000..0f47bf7 --- /dev/null +++ b/services/knowledge-base-service/app/api/v1/__init__.py @@ -0,0 +1,6 @@ +from fastapi import APIRouter + +from app.api.v1.routes import router as routes_router + +router = APIRouter() +router.include_router(routes_router) diff --git a/services/knowledge-base-service/app/api/v1/routes.py b/services/knowledge-base-service/app/api/v1/routes.py new file mode 100644 index 0000000..e767079 --- /dev/null +++ b/services/knowledge-base-service/app/api/v1/routes.py @@ -0,0 +1,286 @@ +from __future__ import annotations + +import json +from typing import Any +from uuid import UUID + +from fastapi import APIRouter, Depends, HTTPException, Request +from qdrant_client.http import models as rest +from sqlalchemy import desc, func, select +from sqlalchemy.orm import Session, joinedload + +from app.config import settings +from app.db import SessionLocal +from app.models import LawArticle, LawDocument +from app.schemas import ( + ArticleResponse, + LawArticleSummary, + LawDocumentSchema, + RetrieveItem, + RetrieveRequest, + RetrieveResponse, + SearchRequest, + SearchResultItem, +) +from app.services.embedding import EmbeddingServiceClient +from redis.asyncio import Redis + +router = APIRouter() + + +def get_db() -> Session: + db = SessionLocal() + try: + yield db + finally: + db.close() + + +def get_qdrant(request: Request): + return request.app.state.qdrant + + +def get_embedding_client(request: Request) -> EmbeddingServiceClient: + return request.app.state.embedding_client + + +def get_redis(request: Request) -> Redis: + return request.app.state.redis + + +@router.get("/health") +async def health(request: Request, db: Session = Depends(get_db)) -> Any: + try: + db.execute(select(1)).scalar_one() + qdrant_client = get_qdrant(request) + qdrant_client.get_collection(settings.qdrant_collection) + except Exception as exc: + raise HTTPException(status_code=503, detail=str(exc)) + + return { + "status": "ok", + "database": "ok", + "qdrant_collection": settings.qdrant_collection, + } + + +def _build_limit(limit: int | None) -> int: + if limit is None or limit <= 0: + return settings.max_search_results + return min(limit, settings.max_search_results) + + +@router.post("/search") +async def search( + payload: SearchRequest, + request: Request, + db: Session = Depends(get_db), + qdrant_client=Depends(get_qdrant), + embedding_client: EmbeddingServiceClient = Depends(get_embedding_client), +) -> list[SearchResultItem]: + if not payload.query.strip(): + raise HTTPException(status_code=400, detail="query is required") + + limit = _build_limit(payload.limit) + vector_results: dict[UUID, dict[str, Any]] = {} + qdrant_filter = None + if payload.domain: + qdrant_filter = rest.Filter( + must=[ + rest.FieldCondition( + key="domain", + match=rest.MatchValue(value=payload.domain), + ) + ] + ) + + embedding = await embedding_client.embed(payload.query) + qdrant_hits = qdrant_client.search( + collection_name=settings.qdrant_collection, + query_vector=embedding, + limit=limit, + with_payload=True, + query_filter=qdrant_filter, + ) + + for rank, hit in enumerate(qdrant_hits, start=1): + payload_data = hit.payload or {} + article_id = payload_data.get("article_id") + if not article_id or article_id in vector_results: + continue + + vector_results[article_id] = { + "article_id": UUID(str(article_id)), + "law_name": payload_data.get("law_name", ""), + "article_number": payload_data.get("article_number", ""), + "title": payload_data.get("title"), + "domain": payload_data.get("domain", ""), + "language": payload_data.get("language", ""), + "text_preview": payload_data.get("text_preview", ""), + "score": 1.0 / (60 + rank), + } + + query_expression = func.plainto_tsquery("simple", payload.query) + statement = ( + select(LawArticle, func.ts_rank_cd(LawArticle.search_vector, query_expression).label("rank")) + .options(joinedload(LawArticle.document)) + .where(LawArticle.search_vector.op("@@")(query_expression)) + ) + if payload.domain: + statement = statement.where(LawArticle.domain == payload.domain) + if payload.language: + statement = statement.where(LawArticle.language == payload.language) + statement = statement.order_by(desc("rank")).limit(limit) + + keyword_rows = db.execute(statement).all() + keyword_results: dict[UUID, dict[str, Any]] = {} + for index, (article, rank) in enumerate(keyword_rows, start=1): + if article.id in keyword_results: + continue + + keyword_results[article.id] = { + "article_id": article.id, + "law_name": article.document.name if article.document else "", + "article_number": article.article_number, + "title": article.title, + "domain": article.domain, + "language": article.language, + "text_preview": article.full_text[:200], + "score": 1.0 / (60 + index), + } + + merged: dict[UUID, dict[str, Any]] = {} + for article_id, entry in vector_results.items(): + merged[article_id] = entry.copy() + + for article_id, entry in keyword_results.items(): + if article_id in merged: + merged[article_id]["score"] += entry["score"] + else: + merged[article_id] = entry.copy() + + ordered = sorted(merged.values(), key=lambda item: item["score"], reverse=True) + return [SearchResultItem(**item) for item in ordered[:limit]] + + +@router.get("/articles/{article_id}") +async def get_article( + article_id: UUID, + request: Request, + db: Session = Depends(get_db), + redis_client: Redis = Depends(get_redis), +) -> ArticleResponse: + cache_key = f"article:{article_id}" + cached = await redis_client.get(cache_key) + if cached: + return ArticleResponse.model_validate_json(cached) + + article = db.get(LawArticle, article_id) + if not article: + raise HTTPException(status_code=404, detail="Article not found") + + if not article.document: + db.refresh(article, attribute_names=["document"]) + + response = ArticleResponse( + id=article.id, + document_id=article.document_id, + document_code=article.document.code if article.document else None, + document_name=article.document.name if article.document else None, + article_number=article.article_number, + chapter=article.chapter, + title=article.title, + full_text=article.full_text, + plain_summary=article.plain_summary, + domain=article.domain, + language=article.language, + qdrant_id=article.qdrant_id, + ) + await redis_client.set( + cache_key, + response.model_dump_json(), + ex=settings.plain_summary_cache_ttl_seconds, + ) + return response + + +@router.get("/articles") +def list_articles( + law_code: str | None = None, + domain: str | None = None, + language: str | None = None, + db: Session = Depends(get_db), +) -> list[LawArticleSummary]: + statement = select(LawArticle).options(joinedload(LawArticle.document)) + if law_code: + statement = statement.join(LawArticle.document).where(LawDocument.code == law_code) + if domain: + statement = statement.where(LawArticle.domain == domain) + if language: + statement = statement.where(LawArticle.language == language) + + articles = db.scalars(statement).all() + return [ + LawArticleSummary( + id=article.id, + document_id=article.document_id, + document_code=article.document.code if article.document else None, + document_name=article.document.name if article.document else None, + article_number=article.article_number, + title=article.title, + domain=article.domain, + language=article.language, + ) + for article in articles + ] + + +@router.get("/laws") +def list_laws(db: Session = Depends(get_db)) -> list[LawDocumentSchema]: + documents = db.scalars(select(LawDocument)).all() + return [LawDocumentSchema.model_validate(document) for document in documents] + + +@router.post("/internal/retrieve") +def internal_retrieve( + payload: RetrieveRequest, + request: Request, + qdrant_client=Depends(get_qdrant), +) -> RetrieveResponse: + limit = _build_limit(payload.top_k) + qdrant_filter = None + if payload.domain_filter: + conditions = [ + rest.FieldCondition(key="domain", match=rest.MatchValue(value=value)) + for value in payload.domain_filter + ] + qdrant_filter = rest.Filter(min_should=rest.MinShould(conditions=conditions, min_count=1)) + + hits = qdrant_client.search( + collection_name=settings.qdrant_collection, + query_vector=payload.query_vector, + limit=limit, + with_payload=True, + query_filter=qdrant_filter, + ) + + results = [] + for hit in hits: + payload_data = hit.payload or {} + article_id_value = payload_data.get("article_id") + if not article_id_value: + continue + + results.append( + RetrieveItem( + article_id=UUID(str(article_id_value)), + law_name=payload_data.get("law_name", ""), + article_number=payload_data.get("article_number", ""), + domain=payload_data.get("domain", ""), + language=payload_data.get("language", ""), + text_preview=payload_data.get("text_preview", ""), + score=hit.score or 0.0, + ) + ) + + return RetrieveResponse(results=results) diff --git a/services/knowledge-base-service/app/config.py b/services/knowledge-base-service/app/config.py new file mode 100644 index 0000000..5939cea --- /dev/null +++ b/services/knowledge-base-service/app/config.py @@ -0,0 +1,32 @@ +from pydantic import Field +from pydantic_settings import BaseSettings, SettingsConfigDict + + +class Settings(BaseSettings): + model_config = SettingsConfigDict( + env_file=".env", + env_ignore_empty=True, + extra="ignore", + ) + + service_name: str = Field("knowledge-base-service", alias="SERVICE_NAME") + api_prefix: str = Field("/api/v1", alias="API_PREFIX") + database_url: str = Field( + "postgresql+psycopg://lexcam:lexcam_dev@postgres:5432/lexcam_knowledge", + alias="DATABASE_URL", + ) + qdrant_url: str = Field("http://qdrant:6333", alias="QDRANT_URL") + qdrant_api_key: str | None = Field(None, alias="QDRANT_API_KEY") + qdrant_collection: str = Field("lexcam_laws", alias="QDRANT_COLLECTION") + embedding_service_url: str = Field( + "http://embedding-service:8000", alias="EMBEDDING_SERVICE_URL" + ) + redis_url: str = Field("redis://redis:6379/0", alias="REDIS_URL") + log_level: str = Field("INFO", alias="LOG_LEVEL") + plain_summary_cache_ttl_seconds: int = Field( + 7 * 24 * 60 * 60, alias="PLAIN_SUMMARY_CACHE_TTL_SECONDS" + ) + max_search_results: int = Field(10, alias="MAX_SEARCH_RESULTS") + + +settings = Settings() diff --git a/services/knowledge-base-service/app/db.py b/services/knowledge-base-service/app/db.py new file mode 100644 index 0000000..9a7fb6c --- /dev/null +++ b/services/knowledge-base-service/app/db.py @@ -0,0 +1,12 @@ +from sqlalchemy import create_engine +from sqlalchemy.orm import sessionmaker + +from app.config import settings + +engine = create_engine( + settings.database_url.replace("postgresql://", "postgresql+psycopg://", 1), + future=True, + echo=False, + pool_pre_ping=True, +) +SessionLocal = sessionmaker(bind=engine, autoflush=False, autocommit=False, future=True) diff --git a/services/knowledge-base-service/app/main.py b/services/knowledge-base-service/app/main.py new file mode 100644 index 0000000..ce0b1fc --- /dev/null +++ b/services/knowledge-base-service/app/main.py @@ -0,0 +1,42 @@ +from __future__ import annotations + +import logging +from contextlib import asynccontextmanager + +from fastapi import FastAPI + +from app.api.v1 import router as v1_router +from app.config import settings +from app.db import engine +from app.models import Base +from app.services.cache import create_redis_client +from app.services.embedding import EmbeddingServiceClient +from app.services.qdrant import create_qdrant_client + + +def _configure_logging() -> None: + logging.basicConfig( + level=settings.log_level, + format="%(asctime)s %(levelname)s %(name)s %(message)s", + ) + + +@asynccontextmanager +async def lifespan(app: FastAPI): + _configure_logging() + Base.metadata.create_all(bind=engine) + app.state.qdrant = create_qdrant_client() + app.state.embedding_client = EmbeddingServiceClient() + app.state.redis = create_redis_client() + yield + await app.state.embedding_client.close() + await app.state.redis.close() + + +app = FastAPI( + title="LexCam Knowledge Base Service", + version="0.1.0", + lifespan=lifespan, +) + +app.include_router(v1_router, prefix=settings.api_prefix) diff --git a/services/knowledge-base-service/app/models.py b/services/knowledge-base-service/app/models.py new file mode 100644 index 0000000..e0d1308 --- /dev/null +++ b/services/knowledge-base-service/app/models.py @@ -0,0 +1,41 @@ +import uuid +from sqlalchemy import Column, DateTime, ForeignKey, Index, String, Text, func +from sqlalchemy.dialects.postgresql import TSVECTOR, UUID +from sqlalchemy.orm import declarative_base, relationship + +Base = declarative_base() + + +class LawDocument(Base): + __tablename__ = "law_documents" + + id = Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4) + code = Column(String(length=128), nullable=False) + name = Column(String(length=256), nullable=False) + jurisdiction = Column(String(length=128), nullable=False) + language = Column(String(length=16), nullable=False) + version = Column(String(length=64), nullable=True) + created_at = Column(DateTime(timezone=True), server_default=func.now()) + + articles = relationship("LawArticle", back_populates="document") + + +class LawArticle(Base): + __tablename__ = "law_articles" + + id = Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4) + document_id = Column(UUID(as_uuid=True), ForeignKey("law_documents.id"), nullable=False) + article_number = Column(String(length=128), nullable=False) + chapter = Column(String(length=128), nullable=True) + title = Column(String(length=512), nullable=True) + full_text = Column(Text, nullable=False) + plain_summary = Column(Text, nullable=True) + domain = Column(String(length=64), nullable=False) + language = Column(String(length=16), nullable=False) + qdrant_id = Column(String(length=128), nullable=True) + search_vector = Column(TSVECTOR, nullable=True) + + document = relationship("LawDocument", back_populates="articles") + + +Index("ix_law_articles_search_vector", LawArticle.search_vector, postgresql_using="gin") diff --git a/services/knowledge-base-service/app/schemas.py b/services/knowledge-base-service/app/schemas.py new file mode 100644 index 0000000..969fc94 --- /dev/null +++ b/services/knowledge-base-service/app/schemas.py @@ -0,0 +1,85 @@ +from datetime import datetime +from typing import List, Optional +from uuid import UUID + +from pydantic import BaseModel, ConfigDict + + +class LawDocumentSchema(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: UUID + code: str + name: str + jurisdiction: str + language: str + version: Optional[str] + created_at: datetime + + +class LawArticleSummary(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: UUID + document_id: UUID + document_code: Optional[str] + document_name: Optional[str] + article_number: str + title: Optional[str] + domain: str + language: str + + +class ArticleResponse(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: UUID + document_id: UUID + document_code: Optional[str] + document_name: Optional[str] + article_number: str + chapter: Optional[str] + title: Optional[str] + full_text: str + plain_summary: Optional[str] + domain: str + language: str + qdrant_id: Optional[str] + + +class SearchRequest(BaseModel): + query: str + language: Optional[str] = None + domain: Optional[str] = None + limit: Optional[int] = 10 + + +class SearchResultItem(BaseModel): + article_id: UUID + law_name: str + article_number: str + title: Optional[str] + domain: str + language: str + text_preview: str + score: float + + +class RetrieveRequest(BaseModel): + query_vector: List[float] + top_k: Optional[int] = 5 + domain_filter: Optional[List[str]] = None + + +class RetrieveItem(BaseModel): + article_id: UUID + law_name: str + article_number: str + domain: str + language: str + text_preview: str + score: float + + +class RetrieveResponse(BaseModel): + results: List[RetrieveItem] diff --git a/services/knowledge-base-service/app/services/cache.py b/services/knowledge-base-service/app/services/cache.py new file mode 100644 index 0000000..9daf2e2 --- /dev/null +++ b/services/knowledge-base-service/app/services/cache.py @@ -0,0 +1,7 @@ +from redis.asyncio import Redis + +from app.config import settings + + +def create_redis_client() -> Redis: + return Redis.from_url(settings.redis_url, decode_responses=True) diff --git a/services/knowledge-base-service/app/services/embedding.py b/services/knowledge-base-service/app/services/embedding.py new file mode 100644 index 0000000..3135e99 --- /dev/null +++ b/services/knowledge-base-service/app/services/embedding.py @@ -0,0 +1,23 @@ +import httpx + +from app.config import settings + + +class EmbeddingServiceClient: + def __init__(self) -> None: + self._client = httpx.AsyncClient(timeout=30.0) + + async def close(self) -> None: + await self._client.aclose() + + async def embed(self, text: str) -> list[float]: + response = await self._client.post( + f"{settings.embedding_service_url}{settings.api_prefix}/embed", + json={"texts": [text]}, + ) + response.raise_for_status() + payload = response.json() + embeddings = payload.get("embeddings") + if not embeddings or not isinstance(embeddings, list): + raise RuntimeError("Invalid embedding response from Embedding Service") + return embeddings[0] diff --git a/services/knowledge-base-service/app/services/qdrant.py b/services/knowledge-base-service/app/services/qdrant.py new file mode 100644 index 0000000..653dd36 --- /dev/null +++ b/services/knowledge-base-service/app/services/qdrant.py @@ -0,0 +1,10 @@ +from qdrant_client import QdrantClient + +from app.config import settings + + +def create_qdrant_client() -> QdrantClient: + client_kwargs = {"url": settings.qdrant_url} + if settings.qdrant_api_key: + client_kwargs["api_key"] = settings.qdrant_api_key + return QdrantClient(**client_kwargs) diff --git a/services/knowledge-base-service/requirements.txt b/services/knowledge-base-service/requirements.txt new file mode 100644 index 0000000..520ec43 --- /dev/null +++ b/services/knowledge-base-service/requirements.txt @@ -0,0 +1,10 @@ +fastapi==0.111.0 +uvicorn[standard]==0.30.1 +gunicorn==22.0.0 +pydantic-settings==2.3.4 +sqlalchemy==2.0.32 +psycopg[binary]==3.3.1 +qdrant-client==1.8.1 +httpx==0.28.1 +redis==5.3.1 +pytest==7.4.0 diff --git a/services/knowledge-base-service/scripts/seed_kb.py b/services/knowledge-base-service/scripts/seed_kb.py new file mode 100644 index 0000000..bda5b39 --- /dev/null +++ b/services/knowledge-base-service/scripts/seed_kb.py @@ -0,0 +1,124 @@ +"""Seed the Knowledge Base with a sample document and article. + +Usage: + python scripts/seed_kb.py + +Requires environment variables (defaults provided for local dev with docker-compose): +- DATABASE_URL +- QDRANT_URL +- EMBEDDING_SERVICE_URL +""" +import os +import uuid +import json + +import httpx +from qdrant_client import QdrantClient +from qdrant_client.http import models as rest +import psycopg + +DATABASE_URL = os.environ.get( + "DATABASE_URL", "postgresql+psycopg://lexcam:lexcam_dev@localhost:5432/lexcam_knowledge" +) +PSYCOPG_DATABASE_URL = DATABASE_URL.replace("postgresql+psycopg://", "postgresql://", 1) +QDRANT_URL = os.environ.get("QDRANT_URL", "http://localhost:6333") +QDRANT_COLLECTION = os.environ.get("QDRANT_COLLECTION", "lexcam_laws") +EMBEDDING_SERVICE_URL = os.environ.get("EMBEDDING_SERVICE_URL", "http://localhost:8001/api/v1") + +SAMPLE_DOC = { + "code": "LABOR", + "name": "Cameroon Labour Code", + "jurisdiction": "cameroon", + "language": "fr", +} + +SAMPLE_ARTICLE = { + "article_number": "Art. 34", + "chapter": None, + "title": "Protection des travailleurs", + "full_text": "Tout employeur est tenu d'assurer la protection de la santé et de la sécurité des travailleurs.", + "plain_summary": "Employers must ensure worker health and safety.", + "domain": "labor", + "language": "fr", +} + + +def embed_text(text: str) -> list[float]: + url = f"{EMBEDDING_SERVICE_URL}/embed" + with httpx.Client(timeout=30.0) as client: + r = client.post(url, json={"texts": [text]}) + r.raise_for_status() + payload = r.json() + embeddings = payload.get("embeddings") + if not embeddings: + raise RuntimeError("No embeddings returned") + return embeddings[0] + + +def upsert_qdrant(client: QdrantClient, point_id: str, vector: list[float], payload: dict): + client.upsert( + collection_name=QDRANT_COLLECTION, + points=[rest.PointStruct(id=point_id, vector=vector, payload=payload)], + ) + + +def main(): + print("Connecting to DB...", PSYCOPG_DATABASE_URL) + with psycopg.connect(PSYCOPG_DATABASE_URL, autocommit=True) as conn: + with conn.cursor() as cur: + # Insert document + cur.execute( + "INSERT INTO law_documents (code, name, jurisdiction, language) VALUES (%s,%s,%s,%s) RETURNING id", + (SAMPLE_DOC["code"], SAMPLE_DOC["name"], SAMPLE_DOC["jurisdiction"], SAMPLE_DOC["language"]), + ) + doc_id = cur.fetchone()[0] + print("Inserted document id", doc_id) + + # Insert article + cur.execute( + "INSERT INTO law_articles (document_id, article_number, chapter, title, full_text, plain_summary, domain, language) VALUES (%s,%s,%s,%s,%s,%s,%s,%s) RETURNING id", + ( + doc_id, + SAMPLE_ARTICLE["article_number"], + SAMPLE_ARTICLE["chapter"], + SAMPLE_ARTICLE["title"], + SAMPLE_ARTICLE["full_text"], + SAMPLE_ARTICLE["plain_summary"], + SAMPLE_ARTICLE["domain"], + SAMPLE_ARTICLE["language"], + ), + ) + article_id = cur.fetchone()[0] + print("Inserted article id", article_id) + + # Get embedding and upsert to Qdrant + print("Requesting embedding...") + vector = embed_text(SAMPLE_ARTICLE["full_text"]) # list of floats + + q = QdrantClient(url=QDRANT_URL) + point_id = str(uuid.uuid4()) + payload = { + "article_id": str(article_id), + "law_name": SAMPLE_DOC["name"], + "article_number": SAMPLE_ARTICLE["article_number"], + "domain": SAMPLE_ARTICLE["domain"], + "language": SAMPLE_ARTICLE["language"], + "text_preview": SAMPLE_ARTICLE["full_text"][:200], + } + print("Upserting to Qdrant, point id", point_id) + upsert_qdrant(q, point_id, vector, payload) + + # Update article qdrant_id in DB + with psycopg.connect(PSYCOPG_DATABASE_URL, autocommit=True) as conn: + with conn.cursor() as cur: + cur.execute( + "UPDATE law_articles SET qdrant_id = %s WHERE id = %s", + (point_id, article_id), + ) + print("Updated article qdrant_id") + + print("Seed complete") + + +if __name__ == "__main__": + main() diff --git a/services/knowledge-base-service/sql/init_kb_schema.sql b/services/knowledge-base-service/sql/init_kb_schema.sql new file mode 100644 index 0000000..4fec0c0 --- /dev/null +++ b/services/knowledge-base-service/sql/init_kb_schema.sql @@ -0,0 +1,42 @@ +-- Knowledge Base schema for lexcam_knowledge +CREATE EXTENSION IF NOT EXISTS pgcrypto; + +CREATE TABLE IF NOT EXISTS law_documents ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + code VARCHAR(128) NOT NULL, + name VARCHAR(256) NOT NULL, + jurisdiction VARCHAR(128) NOT NULL, + language VARCHAR(16) NOT NULL, + version VARCHAR(64), + created_at TIMESTAMPTZ DEFAULT now() +); + +CREATE TABLE IF NOT EXISTS law_articles ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + document_id UUID NOT NULL REFERENCES law_documents(id) ON DELETE CASCADE, + article_number VARCHAR(128) NOT NULL, + chapter VARCHAR(128), + title VARCHAR(512), + full_text TEXT NOT NULL, + plain_summary TEXT, + domain VARCHAR(64) NOT NULL, + language VARCHAR(16) NOT NULL, + qdrant_id VARCHAR(128), + search_vector TSVECTOR +); + +CREATE INDEX IF NOT EXISTS ix_law_articles_search_vector ON law_articles USING GIN(search_vector); + +CREATE OR REPLACE FUNCTION law_articles_search_vector_trigger() RETURNS trigger AS $$ +begin + new.search_vector := to_tsvector('simple', coalesce(new.full_text, '')); + return new; +end +$$ LANGUAGE plpgsql; + +DROP TRIGGER IF EXISTS tsvectorupdate ON law_articles; +CREATE TRIGGER tsvectorupdate BEFORE INSERT OR UPDATE ON law_articles +FOR EACH ROW EXECUTE FUNCTION law_articles_search_vector_trigger(); + +-- Optional helper to reindex all existing rows +-- UPDATE law_articles SET search_vector = to_tsvector('simple', coalesce(full_text, ''));