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} .
${DOCKER} build -f containers/Containerfile.flow \
-t ${CONTAINER_BASE}/trustgraph-flow:${VERSION} .
# ${DOCKER} build -f containers/Containerfile.vertexai \
# -t ${CONTAINER_BASE}/trustgraph-vertexai:${VERSION} .
${DOCKER} build -f containers/Containerfile.vertexai \
-t ${CONTAINER_BASE}/trustgraph-vertexai:${VERSION} .
basic-containers: update-package-versions
${DOCKER} build -f containers/Containerfile.base \

View file

@ -3,7 +3,12 @@
import pulsar
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="""
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 . graph_embeddings_client import GraphEmbeddingsClientSpec
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)
entities = self.query_document_embeddings(request)
docs = await self.query_document_embeddings(request)
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})
print("Done.", flush=True)

View file

@ -53,6 +53,16 @@ class PromptClient(RequestResponse):
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):
def __init__(
self, request_name, response_name,

View file

@ -11,8 +11,6 @@ class Document(Record):
metadata = Metadata()
data = Bytes()
document_ingest_queue = topic('document-load')
############################################################################
# Text documents / text from PDF
@ -21,8 +19,6 @@ class TextDocument(Record):
metadata = Metadata()
text = Bytes()
text_ingest_queue = topic('text-document-load')
############################################################################
# Chunks of text
@ -31,8 +27,6 @@ class Chunk(Record):
metadata = Metadata()
chunk = Bytes()
chunk_ingest_queue = topic('chunk-load')
############################################################################
# Document embeddings are embeddings associated with a chunk
@ -46,8 +40,6 @@ class DocumentEmbeddings(Record):
metadata = Metadata()
chunks = Array(ChunkEmbeddings())
document_embeddings_store_queue = topic('document-embeddings-store')
############################################################################
# Doc embeddings query
@ -62,10 +54,3 @@ class DocumentEmbeddingsResponse(Record):
error = Error()
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 prometheus_client import Histogram
from ... schema import TextDocument, Chunk, Metadata
from ... schema import text_ingest_queue, chunk_ingest_queue
from ... log_level import LogLevel
from ... schema import TextDocument, Chunk
from ... base import FlowProcessor, ConsumerSpec, ProducerSpec
default_ident = "chunker"

View file

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

View file

@ -9,7 +9,6 @@ import base64
from langchain_community.document_loaders import PyPDFLoader
from ... schema import Document, TextDocument, Metadata
from ... schema import document_ingest_queue, text_ingest_queue
from ... log_level import LogLevel
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)
async def handle(self, msg):
async def query_document_embeddings(self, msg):
try:
v = msg.value()
# Sender-produced ID
id = msg.properties()["id"]
print(f"Handling input {id}...", flush=True)
chunks = []
for vec in v.vectors:
for vec in msg.vectors:
dim = len(vec)
collection = (
"d_" + v.user + "_" + v.collection + "_" +
"d_" + msg.user + "_" + msg.collection + "_" +
str(dim)
)
search_result = self.qdrant.query_points(
collection_name=collection,
query=vec,
limit=v.limit,
limit=msg.limit,
with_payload=True,
).points

View file

@ -1,20 +1,7 @@
from . clients.document_embeddings_client import DocumentEmbeddingsClient
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
import asyncio
LABEL="http://www.w3.org/2000/01/rdf-schema#label"
DEFINITION="http://www.w3.org/2004/02/skos/core#definition"
class Query:
@ -28,27 +15,28 @@ class Query:
self.verbose = verbose
self.doc_limit = doc_limit
def get_vector(self, query):
async def get_vector(self, query):
if self.verbose:
print("Compute embeddings...", flush=True)
qembeds = self.rag.embeddings.request(query)
qembeds = await self.rag.embeddings_client.embed(query)
if self.verbose:
print("Done.", flush=True)
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:
print("Get entities...", flush=True)
print("Get docs...", flush=True)
docs = self.rag.de_client.request(
vectors, limit=self.doc_limit
docs = await self.rag.doc_embeddings_client.query(
vectors, limit=self.doc_limit,
user=self.user, collection=self.collection,
)
if self.verbose:
@ -61,70 +49,20 @@ class Query:
class DocumentRag:
def __init__(
self,
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,
self, prompt_client, embeddings_client, doc_embeddings_client,
verbose=False,
module="test",
):
self.verbose=verbose
self.verbose = verbose
if pr_request_queue is None:
pr_request_queue = prompt_request_queue
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,
)
self.prompt_client = prompt_client
self.embeddings_client = embeddings_client
self.doc_embeddings_client = doc_embeddings_client
if self.verbose:
print("Initialised", flush=True)
def query(
async def query(
self, query, user="trustgraph", collection="default",
doc_limit=20,
):
@ -137,14 +75,17 @@ class DocumentRag:
doc_limit=doc_limit
)
docs = q.get_docs(query)
docs = await q.get_docs(query)
if self.verbose:
print("Invoke LLM...", flush=True)
print(docs)
print(query)
resp = self.lang.request_document_prompt(query, docs)
resp = await self.prompt_client.document_prompt(
query = query,
documents = docs
)
if self.verbose:
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 document_rag_request_queue, document_rag_response_queue
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
from ... log_level import LogLevel
from ... document_rag import DocumentRag
from ... base import ConsumerProducer
from . document_rag import DocumentRag
from ... base import FlowProcessor, ConsumerSpec, ProducerSpec
from ... base import PromptClientSpec, EmbeddingsClientSpec
from ... base import DocumentEmbeddingsClientSpec
module = "document-rag"
default_ident = "document-rag"
default_input_queue = document_rag_request_queue
default_output_queue = document_rag_response_queue
default_subscriber = module
class Processor(ConsumerProducer):
class Processor(FlowProcessor):
def __init__(self, **params):
input_queue = params.get("input_queue", default_input_queue)
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
)
id = params.get("id", default_ident)
doc_limit = params.get("doc_limit", 10)
doc_limit = params.get("doc_limit", 5)
super(Processor, self).__init__(
**params | {
"input_queue": input_queue,
"output_queue": output_queue,
"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,
"id": id,
"doc_limit": doc_limit,
}
)
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
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:
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()
# Sender-produced ID
@ -99,11 +88,15 @@ class Processor(ConsumerProducer):
else:
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)
r = DocumentRagResponse(response = response, error=None)
await self.send(r, properties={"id": id})
await flow("response").send(
DocumentRagResponse(
response = response,
error = None
),
properties = {"id": id}
)
print("Done.", flush=True)
@ -113,25 +106,21 @@ class Processor(ConsumerProducer):
print("Send error response...", flush=True)
r = DocumentRagResponse(
error=Error(
type = "llm-error",
message = str(e),
await flow("response").send(
DocumentRagResponse(
response = None,
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
def add_args(parser):
ConsumerProducer.add_args(
parser, default_input_queue, default_subscriber,
default_output_queue,
)
FlowProcessor.add_args(parser)
parser.add_argument(
'-d', '--doc-limit',
@ -140,43 +129,7 @@ class Processor(ConsumerProducer):
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():
Processor.launch(module, __doc__)
Processor.launch(default_ident, __doc__)

View file

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