shak3008 commited on
Commit
93def79
·
1 Parent(s): af21975

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", "classify")
 
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