Add document chunk embeddings to workflow
Browse files
backend/alembic/versions/585aa57edb27_add_chunk_embeddings.py
ADDED
|
@@ -0,0 +1,30 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""add chunk embeddings
|
| 2 |
+
|
| 3 |
+
Revision ID: 585aa57edb27
|
| 4 |
+
Revises: 5b7572e0875a
|
| 5 |
+
Create Date: 2026-08-10 22:06:30.806027
|
| 6 |
+
|
| 7 |
+
"""
|
| 8 |
+
from typing import Sequence, Union
|
| 9 |
+
|
| 10 |
+
from alembic import op
|
| 11 |
+
import sqlalchemy as sa
|
| 12 |
+
|
| 13 |
+
from pgvector.sqlalchemy import Vector
|
| 14 |
+
|
| 15 |
+
revision: str = '585aa57edb27'
|
| 16 |
+
down_revision: Union[str, None] = '5b7572e0875a'
|
| 17 |
+
branch_labels: Union[str, Sequence[str], None] = None
|
| 18 |
+
depends_on: Union[str, Sequence[str], None] = None
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
def upgrade() -> None:
|
| 22 |
+
# ### commands auto generated by Alembic - please adjust! ###
|
| 23 |
+
op.add_column('document_chunks', sa.Column('embedding', Vector(384), nullable=True))
|
| 24 |
+
# ### end Alembic commands ###
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def downgrade() -> None:
|
| 28 |
+
# ### commands auto generated by Alembic - please adjust! ###
|
| 29 |
+
op.drop_column('document_chunks', 'embedding')
|
| 30 |
+
# ### end Alembic commands ###
|
backend/app/models/document_chunk.py
CHANGED
|
@@ -8,7 +8,7 @@ from sqlalchemy.dialects.postgresql import UUID
|
|
| 8 |
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
| 9 |
|
| 10 |
from app.database.database import Base
|
| 11 |
-
|
| 12 |
|
| 13 |
class DocumentChunk(Base):
|
| 14 |
"""
|
|
@@ -56,6 +56,11 @@ class DocumentChunk(Base):
|
|
| 56 |
nullable=True,
|
| 57 |
)
|
| 58 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 59 |
created_at: Mapped[datetime] = mapped_column(
|
| 60 |
DateTime(timezone=True),
|
| 61 |
server_default=func.now(),
|
|
|
|
| 8 |
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
| 9 |
|
| 10 |
from app.database.database import Base
|
| 11 |
+
from pgvector.sqlalchemy import Vector
|
| 12 |
|
| 13 |
class DocumentChunk(Base):
|
| 14 |
"""
|
|
|
|
| 56 |
nullable=True,
|
| 57 |
)
|
| 58 |
|
| 59 |
+
embedding: Mapped[list[float] | None] = mapped_column(
|
| 60 |
+
Vector(384),
|
| 61 |
+
nullable=True,
|
| 62 |
+
)
|
| 63 |
+
|
| 64 |
created_at: Mapped[datetime] = mapped_column(
|
| 65 |
DateTime(timezone=True),
|
| 66 |
server_default=func.now(),
|
backend/app/services/embedding_service.py
ADDED
|
@@ -0,0 +1,40 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from sentence_transformers import SentenceTransformer
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
class EmbeddingService:
|
| 7 |
+
"""
|
| 8 |
+
Generates dense vector embeddings for document chunks.
|
| 9 |
+
|
| 10 |
+
Model output dimension:
|
| 11 |
+
384
|
| 12 |
+
"""
|
| 13 |
+
|
| 14 |
+
MODEL_NAME = "sentence-transformers/all-MiniLM-L6-v2"
|
| 15 |
+
DIMENSION = 384
|
| 16 |
+
|
| 17 |
+
def __init__(self):
|
| 18 |
+
self.model = SentenceTransformer(self.MODEL_NAME)
|
| 19 |
+
|
| 20 |
+
def embed(self, text: str) -> list[float]:
|
| 21 |
+
if not text or not text.strip():
|
| 22 |
+
return []
|
| 23 |
+
|
| 24 |
+
vector = self.model.encode(
|
| 25 |
+
text,
|
| 26 |
+
normalize_embeddings=True,
|
| 27 |
+
)
|
| 28 |
+
|
| 29 |
+
return vector.tolist()
|
| 30 |
+
|
| 31 |
+
def embed_many(self, texts: list[str]) -> list[list[float]]:
|
| 32 |
+
if not texts:
|
| 33 |
+
return []
|
| 34 |
+
|
| 35 |
+
vectors = self.model.encode(
|
| 36 |
+
texts,
|
| 37 |
+
normalize_embeddings=True,
|
| 38 |
+
)
|
| 39 |
+
|
| 40 |
+
return vectors.tolist()
|
backend/app/workflow/graph.py
CHANGED
|
@@ -16,7 +16,8 @@ from app.workflow.nodes.knowledge import knowledge
|
|
| 16 |
from app.workflow.nodes.link import link
|
| 17 |
from app.workflow.nodes.reconcile import reconcile
|
| 18 |
from app.workflow.state import WorkflowState
|
| 19 |
-
|
|
|
|
| 20 |
|
| 21 |
def build_workflow(db: Session):
|
| 22 |
"""
|
|
@@ -27,7 +28,7 @@ def build_workflow(db: Session):
|
|
| 27 |
knowledge_repository = KnowledgeRepository(db)
|
| 28 |
chunking_service = ChunkingService(db)
|
| 29 |
knowledge_link_service = KnowledgeLinkService(db)
|
| 30 |
-
|
| 31 |
graph = StateGraph(WorkflowState)
|
| 32 |
|
| 33 |
# ------------------------------------------------------------------
|
|
@@ -73,13 +74,22 @@ def build_workflow(db: Session):
|
|
| 73 |
|
| 74 |
graph.add_node("complete", complete)
|
| 75 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 76 |
# ------------------------------------------------------------------
|
| 77 |
# Workflow edges
|
| 78 |
# ------------------------------------------------------------------
|
| 79 |
|
| 80 |
graph.add_edge(START, "extract")
|
| 81 |
graph.add_edge("extract", "chunk")
|
| 82 |
-
graph.add_edge("chunk", "
|
|
|
|
| 83 |
graph.add_edge("classify", "knowledge")
|
| 84 |
graph.add_edge("knowledge", "reconciliation")
|
| 85 |
graph.add_edge("reconciliation", "link")
|
|
|
|
| 16 |
from app.workflow.nodes.link import link
|
| 17 |
from app.workflow.nodes.reconcile import reconcile
|
| 18 |
from app.workflow.state import WorkflowState
|
| 19 |
+
from app.services.embedding_service import EmbeddingService
|
| 20 |
+
from app.workflow.nodes.embed import embed
|
| 21 |
|
| 22 |
def build_workflow(db: Session):
|
| 23 |
"""
|
|
|
|
| 28 |
knowledge_repository = KnowledgeRepository(db)
|
| 29 |
chunking_service = ChunkingService(db)
|
| 30 |
knowledge_link_service = KnowledgeLinkService(db)
|
| 31 |
+
embedding_service = EmbeddingService()
|
| 32 |
graph = StateGraph(WorkflowState)
|
| 33 |
|
| 34 |
# ------------------------------------------------------------------
|
|
|
|
| 74 |
|
| 75 |
graph.add_node("complete", complete)
|
| 76 |
|
| 77 |
+
graph.add_node(
|
| 78 |
+
"embedding",
|
| 79 |
+
lambda state: embed(
|
| 80 |
+
state,
|
| 81 |
+
db,
|
| 82 |
+
embedding_service,
|
| 83 |
+
),
|
| 84 |
+
)
|
| 85 |
# ------------------------------------------------------------------
|
| 86 |
# Workflow edges
|
| 87 |
# ------------------------------------------------------------------
|
| 88 |
|
| 89 |
graph.add_edge(START, "extract")
|
| 90 |
graph.add_edge("extract", "chunk")
|
| 91 |
+
graph.add_edge("chunk", "embedding")
|
| 92 |
+
graph.add_edge("embedding", "classify")
|
| 93 |
graph.add_edge("classify", "knowledge")
|
| 94 |
graph.add_edge("knowledge", "reconciliation")
|
| 95 |
graph.add_edge("reconciliation", "link")
|
backend/app/workflow/nodes/embed.py
ADDED
|
@@ -0,0 +1,52 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from sqlalchemy.orm import Session
|
| 4 |
+
|
| 5 |
+
from app.models.document_chunk import DocumentChunk
|
| 6 |
+
from app.services.embedding_service import EmbeddingService
|
| 7 |
+
from app.workflow.state import WorkflowState
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
def embed(
|
| 11 |
+
state: WorkflowState,
|
| 12 |
+
db: Session,
|
| 13 |
+
embedding_service: EmbeddingService,
|
| 14 |
+
) -> WorkflowState:
|
| 15 |
+
"""
|
| 16 |
+
Generate and persist embeddings for the current document's chunks.
|
| 17 |
+
"""
|
| 18 |
+
|
| 19 |
+
state.current_node = "EMBEDDING"
|
| 20 |
+
|
| 21 |
+
chunks = (
|
| 22 |
+
db.query(DocumentChunk)
|
| 23 |
+
.filter(
|
| 24 |
+
DocumentChunk.document_version_id
|
| 25 |
+
== state.document_version_id,
|
| 26 |
+
DocumentChunk.embedding.is_(None),
|
| 27 |
+
)
|
| 28 |
+
.order_by(DocumentChunk.chunk_index)
|
| 29 |
+
.all()
|
| 30 |
+
)
|
| 31 |
+
|
| 32 |
+
if not chunks:
|
| 33 |
+
state.metadata["embeddings_created"] = 0
|
| 34 |
+
return state
|
| 35 |
+
|
| 36 |
+
texts = [chunk.text for chunk in chunks]
|
| 37 |
+
|
| 38 |
+
vectors = embedding_service.embed_many(texts)
|
| 39 |
+
|
| 40 |
+
if len(vectors) != len(chunks):
|
| 41 |
+
raise RuntimeError(
|
| 42 |
+
"Embedding count does not match chunk count."
|
| 43 |
+
)
|
| 44 |
+
|
| 45 |
+
for chunk, vector in zip(chunks, vectors):
|
| 46 |
+
chunk.embedding = vector
|
| 47 |
+
|
| 48 |
+
db.commit()
|
| 49 |
+
|
| 50 |
+
state.metadata["embeddings_created"] = len(chunks)
|
| 51 |
+
|
| 52 |
+
return state
|
backend/requirements.txt
CHANGED
|
@@ -49,4 +49,7 @@ email-validator==2.2.0
|
|
| 49 |
# Agentic workflow
|
| 50 |
langgraph
|
| 51 |
langchain
|
| 52 |
-
langchain-groq
|
|
|
|
|
|
|
|
|
|
|
|
| 49 |
# Agentic workflow
|
| 50 |
langgraph
|
| 51 |
langchain
|
| 52 |
+
langchain-groq
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
grandalf==0.8
|