Doc RAG working

This commit is contained in:
Cyber MacGeddon 2025-04-22 10:52:11 +01:00
parent 6be7b30633
commit c9913297d2
14 changed files with 157 additions and 237 deletions

View file

@ -65,8 +65,8 @@ some-containers:
-t ${CONTAINER_BASE}/trustgraph-base:${VERSION} . -t ${CONTAINER_BASE}/trustgraph-base:${VERSION} .
${DOCKER} build -f containers/Containerfile.flow \ ${DOCKER} build -f containers/Containerfile.flow \
-t ${CONTAINER_BASE}/trustgraph-flow:${VERSION} . -t ${CONTAINER_BASE}/trustgraph-flow:${VERSION} .
# ${DOCKER} build -f containers/Containerfile.vertexai \ ${DOCKER} build -f containers/Containerfile.vertexai \
# -t ${CONTAINER_BASE}/trustgraph-vertexai:${VERSION} . -t ${CONTAINER_BASE}/trustgraph-vertexai:${VERSION} .
basic-containers: update-package-versions basic-containers: update-package-versions
${DOCKER} build -f containers/Containerfile.base \ ${DOCKER} build -f containers/Containerfile.base \

View file

@ -3,7 +3,12 @@
import pulsar import pulsar
from trustgraph.clients.document_rag_client import DocumentRagClient from trustgraph.clients.document_rag_client import DocumentRagClient
rag = DocumentRagClient(pulsar_host="pulsar://localhost:6650") rag = DocumentRagClient(
pulsar_host="pulsar://localhost:6650",
subscriber="test1",
input_queue = "non-persistent://tg/request/document-rag:default",
output_queue = "non-persistent://tg/response/document-rag:default",
)
query=""" query="""
What was the cause of the space shuttle disaster?""" What was the cause of the space shuttle disaster?"""

View file

@ -25,4 +25,5 @@ from . graph_embeddings_query_service import GraphEmbeddingsQueryService
from . document_embeddings_query_service import DocumentEmbeddingsQueryService from . document_embeddings_query_service import DocumentEmbeddingsQueryService
from . graph_embeddings_client import GraphEmbeddingsClientSpec from . graph_embeddings_client import GraphEmbeddingsClientSpec
from . triples_client import TriplesClientSpec from . triples_client import TriplesClientSpec
from . document_embeddings_client import DocumentEmbeddingsClientSpec

View file

@ -0,0 +1,38 @@
from . request_response_spec import RequestResponse, RequestResponseSpec
from .. schema import DocumentEmbeddingsRequest, DocumentEmbeddingsResponse
from .. knowledge import Uri, Literal
class DocumentEmbeddingsClient(RequestResponse):
async def query(self, vectors, limit=20, user="trustgraph",
collection="default", timeout=30):
resp = await self.request(
DocumentEmbeddingsRequest(
vectors = vectors,
limit = limit,
user = user,
collection = collection
),
timeout=timeout
)
print(resp, flush=True)
if resp.error:
raise RuntimeError(resp.error.message)
return resp.documents
class DocumentEmbeddingsClientSpec(RequestResponseSpec):
def __init__(
self, request_name, response_name,
):
super(DocumentEmbeddingsClientSpec, self).__init__(
request_name = request_name,
request_schema = DocumentEmbeddingsRequest,
response_name = response_name,
response_schema = DocumentEmbeddingsResponse,
impl = DocumentEmbeddingsClient,
)

View file

@ -49,10 +49,10 @@ class DocumentEmbeddingsQueryService(FlowProcessor):
print(f"Handling input {id}...", flush=True) print(f"Handling input {id}...", flush=True)
entities = self.query_document_embeddings(request) docs = await self.query_document_embeddings(request)
print("Send response...", flush=True) print("Send response...", flush=True)
r = DocumentEmbeddingsResponse(entities=entities, error=None) r = DocumentEmbeddingsResponse(documents=docs, error=None)
await flow("response").send(r, properties={"id": id}) await flow("response").send(r, properties={"id": id})
print("Done.", flush=True) print("Done.", flush=True)

View file

@ -53,6 +53,16 @@ class PromptClient(RequestResponse):
timeout = timeout, timeout = timeout,
) )
async def document_prompt(self, query, documents, timeout=600):
return await self.prompt(
id = "document-prompt",
variables = {
"query": query,
"documents": documents,
},
timeout = timeout,
)
class PromptClientSpec(RequestResponseSpec): class PromptClientSpec(RequestResponseSpec):
def __init__( def __init__(
self, request_name, response_name, self, request_name, response_name,

View file

@ -11,8 +11,6 @@ class Document(Record):
metadata = Metadata() metadata = Metadata()
data = Bytes() data = Bytes()
document_ingest_queue = topic('document-load')
############################################################################ ############################################################################
# Text documents / text from PDF # Text documents / text from PDF
@ -21,8 +19,6 @@ class TextDocument(Record):
metadata = Metadata() metadata = Metadata()
text = Bytes() text = Bytes()
text_ingest_queue = topic('text-document-load')
############################################################################ ############################################################################
# Chunks of text # Chunks of text
@ -31,8 +27,6 @@ class Chunk(Record):
metadata = Metadata() metadata = Metadata()
chunk = Bytes() chunk = Bytes()
chunk_ingest_queue = topic('chunk-load')
############################################################################ ############################################################################
# Document embeddings are embeddings associated with a chunk # Document embeddings are embeddings associated with a chunk
@ -46,8 +40,6 @@ class DocumentEmbeddings(Record):
metadata = Metadata() metadata = Metadata()
chunks = Array(ChunkEmbeddings()) chunks = Array(ChunkEmbeddings())
document_embeddings_store_queue = topic('document-embeddings-store')
############################################################################ ############################################################################
# Doc embeddings query # Doc embeddings query
@ -62,10 +54,3 @@ class DocumentEmbeddingsResponse(Record):
error = Error() error = Error()
documents = Array(Bytes()) documents = Array(Bytes())
document_embeddings_request_queue = topic(
'doc-embeddings', kind='non-persistent', namespace='request'
)
document_embeddings_response_queue = topic(
'doc-embeddings', kind='non-persistent', namespace='response',
)

View file

@ -7,9 +7,7 @@ as text as separate output objects.
from langchain_text_splitters import RecursiveCharacterTextSplitter from langchain_text_splitters import RecursiveCharacterTextSplitter
from prometheus_client import Histogram from prometheus_client import Histogram
from ... schema import TextDocument, Chunk, Metadata from ... schema import TextDocument, Chunk
from ... schema import text_ingest_queue, chunk_ingest_queue
from ... log_level import LogLevel
from ... base import FlowProcessor, ConsumerSpec, ProducerSpec from ... base import FlowProcessor, ConsumerSpec, ProducerSpec
default_ident = "chunker" default_ident = "chunker"

View file

@ -7,9 +7,7 @@ as text as separate output objects.
from langchain_text_splitters import TokenTextSplitter from langchain_text_splitters import TokenTextSplitter
from prometheus_client import Histogram from prometheus_client import Histogram
from ... schema import TextDocument, Chunk, Metadata from ... schema import TextDocument, Chunk
from ... schema import text_ingest_queue, chunk_ingest_queue
from ... log_level import LogLevel
from ... base import FlowProcessor from ... base import FlowProcessor
default_ident = "chunker" default_ident = "chunker"

View file

@ -9,7 +9,6 @@ import base64
from langchain_community.document_loaders import PyPDFLoader from langchain_community.document_loaders import PyPDFLoader
from ... schema import Document, TextDocument, Metadata from ... schema import Document, TextDocument, Metadata
from ... schema import document_ingest_queue, text_ingest_queue
from ... log_level import LogLevel from ... log_level import LogLevel
from ... base import FlowProcessor, ConsumerSpec, ProducerSpec from ... base import FlowProcessor, ConsumerSpec, ProducerSpec

View file

@ -34,31 +34,24 @@ class Processor(DocumentEmbeddingsQueryService):
self.qdrant = QdrantClient(url=store_uri, api_key=api_key) self.qdrant = QdrantClient(url=store_uri, api_key=api_key)
async def handle(self, msg): async def query_document_embeddings(self, msg):
try: try:
v = msg.value()
# Sender-produced ID
id = msg.properties()["id"]
print(f"Handling input {id}...", flush=True)
chunks = [] chunks = []
for vec in v.vectors: for vec in msg.vectors:
dim = len(vec) dim = len(vec)
collection = ( collection = (
"d_" + v.user + "_" + v.collection + "_" + "d_" + msg.user + "_" + msg.collection + "_" +
str(dim) str(dim)
) )
search_result = self.qdrant.query_points( search_result = self.qdrant.query_points(
collection_name=collection, collection_name=collection,
query=vec, query=vec,
limit=v.limit, limit=msg.limit,
with_payload=True, with_payload=True,
).points ).points

View file

@ -1,20 +1,7 @@
from . clients.document_embeddings_client import DocumentEmbeddingsClient import asyncio
from . clients.triples_query_client import TriplesQueryClient
from . clients.embeddings_client import EmbeddingsClient
from . clients.prompt_client import PromptClient
from . schema import DocumentEmbeddingsRequest, DocumentEmbeddingsResponse
from . schema import TriplesQueryRequest, TriplesQueryResponse
from . schema import prompt_request_queue
from . schema import prompt_response_queue
from . schema import embeddings_request_queue
from . schema import embeddings_response_queue
from . schema import document_embeddings_request_queue
from . schema import document_embeddings_response_queue
LABEL="http://www.w3.org/2000/01/rdf-schema#label" LABEL="http://www.w3.org/2000/01/rdf-schema#label"
DEFINITION="http://www.w3.org/2004/02/skos/core#definition"
class Query: class Query:
@ -28,27 +15,28 @@ class Query:
self.verbose = verbose self.verbose = verbose
self.doc_limit = doc_limit self.doc_limit = doc_limit
def get_vector(self, query): async def get_vector(self, query):
if self.verbose: if self.verbose:
print("Compute embeddings...", flush=True) print("Compute embeddings...", flush=True)
qembeds = self.rag.embeddings.request(query) qembeds = await self.rag.embeddings_client.embed(query)
if self.verbose: if self.verbose:
print("Done.", flush=True) print("Done.", flush=True)
return qembeds return qembeds
def get_docs(self, query): async def get_docs(self, query):
vectors = self.get_vector(query) vectors = await self.get_vector(query)
if self.verbose: if self.verbose:
print("Get entities...", flush=True) print("Get docs...", flush=True)
docs = self.rag.de_client.request( docs = await self.rag.doc_embeddings_client.query(
vectors, limit=self.doc_limit vectors, limit=self.doc_limit,
user=self.user, collection=self.collection,
) )
if self.verbose: if self.verbose:
@ -61,70 +49,20 @@ class Query:
class DocumentRag: class DocumentRag:
def __init__( def __init__(
self, self, prompt_client, embeddings_client, doc_embeddings_client,
pulsar_host="pulsar://pulsar:6650",
pulsar_api_key=None,
pr_request_queue=None,
pr_response_queue=None,
emb_request_queue=None,
emb_response_queue=None,
de_request_queue=None,
de_response_queue=None,
verbose=False, verbose=False,
module="test",
): ):
self.verbose=verbose self.verbose = verbose
if pr_request_queue is None: self.prompt_client = prompt_client
pr_request_queue = prompt_request_queue self.embeddings_client = embeddings_client
self.doc_embeddings_client = doc_embeddings_client
if pr_response_queue is None:
pr_response_queue = prompt_response_queue
if emb_request_queue is None:
emb_request_queue = embeddings_request_queue
if emb_response_queue is None:
emb_response_queue = embeddings_response_queue
if de_request_queue is None:
de_request_queue = document_embeddings_request_queue
if de_response_queue is None:
de_response_queue = document_embeddings_response_queue
if self.verbose:
print("Initialising...", flush=True)
self.de_client = DocumentEmbeddingsClient(
pulsar_host=pulsar_host,
subscriber=module + "-de",
input_queue=de_request_queue,
output_queue=de_response_queue,
pulsar_api_key=pulsar_api_key,
)
self.embeddings = EmbeddingsClient(
pulsar_host=pulsar_host,
input_queue=emb_request_queue,
output_queue=emb_response_queue,
subscriber=module + "-emb",
pulsar_api_key=pulsar_api_key,
)
self.lang = PromptClient(
pulsar_host=pulsar_host,
input_queue=pr_request_queue,
output_queue=pr_response_queue,
subscriber=module + "-de-prompt",
pulsar_api_key=pulsar_api_key,
)
if self.verbose: if self.verbose:
print("Initialised", flush=True) print("Initialised", flush=True)
def query( async def query(
self, query, user="trustgraph", collection="default", self, query, user="trustgraph", collection="default",
doc_limit=20, doc_limit=20,
): ):
@ -137,14 +75,17 @@ class DocumentRag:
doc_limit=doc_limit doc_limit=doc_limit
) )
docs = q.get_docs(query) docs = await q.get_docs(query)
if self.verbose: if self.verbose:
print("Invoke LLM...", flush=True) print("Invoke LLM...", flush=True)
print(docs) print(docs)
print(query) print(query)
resp = self.lang.request_document_prompt(query, docs) resp = await self.prompt_client.document_prompt(
query = query,
documents = docs
)
if self.verbose: if self.verbose:
print("Done", flush=True) print("Done", flush=True)

View file

@ -5,88 +5,77 @@ Input is query, output is response.
""" """
from ... schema import DocumentRagQuery, DocumentRagResponse, Error from ... schema import DocumentRagQuery, DocumentRagResponse, Error
from ... schema import document_rag_request_queue, document_rag_response_queue from . document_rag import DocumentRag
from ... schema import prompt_request_queue from ... base import FlowProcessor, ConsumerSpec, ProducerSpec
from ... schema import prompt_response_queue from ... base import PromptClientSpec, EmbeddingsClientSpec
from ... schema import embeddings_request_queue from ... base import DocumentEmbeddingsClientSpec
from ... schema import embeddings_response_queue
from ... schema import document_embeddings_request_queue
from ... schema import document_embeddings_response_queue
from ... log_level import LogLevel
from ... document_rag import DocumentRag
from ... base import ConsumerProducer
module = "document-rag" default_ident = "document-rag"
default_input_queue = document_rag_request_queue class Processor(FlowProcessor):
default_output_queue = document_rag_response_queue
default_subscriber = module
class Processor(ConsumerProducer):
def __init__(self, **params): def __init__(self, **params):
input_queue = params.get("input_queue", default_input_queue) id = params.get("id", default_ident)
output_queue = params.get("output_queue", default_output_queue)
subscriber = params.get("subscriber", default_subscriber)
pr_request_queue = params.get(
"prompt_request_queue", prompt_request_queue
)
pr_response_queue = params.get(
"prompt_response_queue", prompt_response_queue
)
emb_request_queue = params.get(
"embeddings_request_queue", embeddings_request_queue
)
emb_response_queue = params.get(
"embeddings_response_queue", embeddings_response_queue
)
de_request_queue = params.get(
"document_embeddings_request_queue",
document_embeddings_request_queue
)
de_response_queue = params.get(
"document_embeddings_response_queue",
document_embeddings_response_queue
)
doc_limit = params.get("doc_limit", 10) doc_limit = params.get("doc_limit", 5)
super(Processor, self).__init__( super(Processor, self).__init__(
**params | { **params | {
"input_queue": input_queue, "id": id,
"output_queue": output_queue, "doc_limit": doc_limit,
"subscriber": subscriber,
"input_schema": DocumentRagQuery,
"output_schema": DocumentRagResponse,
"prompt_request_queue": pr_request_queue,
"prompt_response_queue": pr_response_queue,
"embeddings_request_queue": emb_request_queue,
"embeddings_response_queue": emb_response_queue,
"document_embeddings_request_queue": de_request_queue,
"document_embeddings_response_queue": de_response_queue,
} }
) )
self.rag = DocumentRag(
pulsar_host=self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
pr_request_queue=pr_request_queue,
pr_response_queue=pr_response_queue,
emb_request_queue=emb_request_queue,
emb_response_queue=emb_response_queue,
de_request_queue=de_request_queue,
de_response_queue=de_response_queue,
verbose=True,
module=module,
)
self.doc_limit = doc_limit self.doc_limit = doc_limit
async def handle(self, msg): self.register_specification(
ConsumerSpec(
name = "request",
schema = DocumentRagQuery,
handler = self.on_request,
)
)
self.register_specification(
EmbeddingsClientSpec(
request_name = "embeddings-request",
response_name = "embeddings-response",
)
)
self.register_specification(
DocumentEmbeddingsClientSpec(
request_name = "document-embeddings-request",
response_name = "document-embeddings-response",
)
)
self.register_specification(
PromptClientSpec(
request_name = "prompt-request",
response_name = "prompt-response",
)
)
self.register_specification(
ProducerSpec(
name = "response",
schema = DocumentRagResponse,
)
)
async def on_request(self, msg, consumer, flow):
try: try:
self.rag = DocumentRag(
embeddings_client = flow("embeddings-request"),
doc_embeddings_client = flow("document-embeddings-request"),
prompt_client = flow("prompt-request"),
verbose=True,
)
v = msg.value() v = msg.value()
# Sender-produced ID # Sender-produced ID
@ -99,11 +88,15 @@ class Processor(ConsumerProducer):
else: else:
doc_limit = self.doc_limit doc_limit = self.doc_limit
response = self.rag.query(v.query, doc_limit=doc_limit) response = await self.rag.query(v.query, doc_limit=doc_limit)
print("Send response...", flush=True) await flow("response").send(
r = DocumentRagResponse(response = response, error=None) DocumentRagResponse(
await self.send(r, properties={"id": id}) response = response,
error = None
),
properties = {"id": id}
)
print("Done.", flush=True) print("Done.", flush=True)
@ -113,25 +106,21 @@ class Processor(ConsumerProducer):
print("Send error response...", flush=True) print("Send error response...", flush=True)
r = DocumentRagResponse( await flow("response").send(
error=Error( DocumentRagResponse(
type = "llm-error", response = None,
message = str(e), error = Error(
type = "document-rag-error",
message = str(e),
),
), ),
response=None, properties = {"id": id}
) )
await self.send(r, properties={"id": id})
self.consumer.acknowledge(msg)
@staticmethod @staticmethod
def add_args(parser): def add_args(parser):
ConsumerProducer.add_args( FlowProcessor.add_args(parser)
parser, default_input_queue, default_subscriber,
default_output_queue,
)
parser.add_argument( parser.add_argument(
'-d', '--doc-limit', '-d', '--doc-limit',
@ -140,43 +129,7 @@ class Processor(ConsumerProducer):
help=f'Default document fetch limit (default: 10)' help=f'Default document fetch limit (default: 10)'
) )
parser.add_argument(
'--prompt-request-queue',
default=prompt_request_queue,
help=f'Prompt request queue (default: {prompt_request_queue})',
)
parser.add_argument(
'--prompt-response-queue',
default=prompt_response_queue,
help=f'Prompt response queue (default: {prompt_response_queue})',
)
parser.add_argument(
'--embeddings-request-queue',
default=embeddings_request_queue,
help=f'Embeddings request queue (default: {embeddings_request_queue})',
)
parser.add_argument(
'--embeddings-response-queue',
default=embeddings_response_queue,
help=f'Embeddings response queue (default: {embeddings_response_queue})',
)
parser.add_argument(
'--document-embeddings-request-queue',
default=document_embeddings_request_queue,
help=f'Document embeddings request queue (default: {document_embeddings_request_queue})',
)
parser.add_argument(
'--document-embeddings-response-queue',
default=document_embeddings_response_queue,
help=f'Document embeddings response queue (default: {document_embeddings_response_queue})',
)
def run(): def run():
Processor.launch(module, __doc__) Processor.launch(default_ident, __doc__)

View file

@ -2,7 +2,6 @@
import asyncio import asyncio
LABEL="http://www.w3.org/2000/01/rdf-schema#label" LABEL="http://www.w3.org/2000/01/rdf-schema#label"
DEFINITION="http://www.w3.org/2004/02/skos/core#definition"
class Query: class Query: