diff --git a/Makefile b/Makefile index 73edcf2f..1899e602 100644 --- a/Makefile +++ b/Makefile @@ -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 \ diff --git a/tests/test-doc-rag b/tests/test-doc-rag index 718157b6..b7382bf5 100755 --- a/tests/test-doc-rag +++ b/tests/test-doc-rag @@ -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?""" diff --git a/trustgraph-base/trustgraph/base/__init__.py b/trustgraph-base/trustgraph/base/__init__.py index 23a099a2..ac3bb766 100644 --- a/trustgraph-base/trustgraph/base/__init__.py +++ b/trustgraph-base/trustgraph/base/__init__.py @@ -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 diff --git a/trustgraph-base/trustgraph/base/document_embeddings_client.py b/trustgraph-base/trustgraph/base/document_embeddings_client.py new file mode 100644 index 00000000..86370c52 --- /dev/null +++ b/trustgraph-base/trustgraph/base/document_embeddings_client.py @@ -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, + ) + diff --git a/trustgraph-base/trustgraph/base/document_embeddings_query_service.py b/trustgraph-base/trustgraph/base/document_embeddings_query_service.py index d84def88..0dee7001 100755 --- a/trustgraph-base/trustgraph/base/document_embeddings_query_service.py +++ b/trustgraph-base/trustgraph/base/document_embeddings_query_service.py @@ -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) diff --git a/trustgraph-base/trustgraph/base/prompt_client.py b/trustgraph-base/trustgraph/base/prompt_client.py index 6758bff2..88e1a15f 100644 --- a/trustgraph-base/trustgraph/base/prompt_client.py +++ b/trustgraph-base/trustgraph/base/prompt_client.py @@ -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, diff --git a/trustgraph-base/trustgraph/schema/documents.py b/trustgraph-base/trustgraph/schema/documents.py index fd0049ee..e479371d 100644 --- a/trustgraph-base/trustgraph/schema/documents.py +++ b/trustgraph-base/trustgraph/schema/documents.py @@ -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', -) - diff --git a/trustgraph-flow/trustgraph/chunking/recursive/chunker.py b/trustgraph-flow/trustgraph/chunking/recursive/chunker.py index 6540d0b0..aa48cc57 100755 --- a/trustgraph-flow/trustgraph/chunking/recursive/chunker.py +++ b/trustgraph-flow/trustgraph/chunking/recursive/chunker.py @@ -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" diff --git a/trustgraph-flow/trustgraph/chunking/token/chunker.py b/trustgraph-flow/trustgraph/chunking/token/chunker.py index 5ae0f4f4..ff217350 100755 --- a/trustgraph-flow/trustgraph/chunking/token/chunker.py +++ b/trustgraph-flow/trustgraph/chunking/token/chunker.py @@ -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" diff --git a/trustgraph-flow/trustgraph/decoding/pdf/pdf_decoder.py b/trustgraph-flow/trustgraph/decoding/pdf/pdf_decoder.py index c9b2a4d6..d0669a59 100755 --- a/trustgraph-flow/trustgraph/decoding/pdf/pdf_decoder.py +++ b/trustgraph-flow/trustgraph/decoding/pdf/pdf_decoder.py @@ -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 diff --git a/trustgraph-flow/trustgraph/query/doc_embeddings/qdrant/service.py b/trustgraph-flow/trustgraph/query/doc_embeddings/qdrant/service.py index 83cfc766..c5543690 100755 --- a/trustgraph-flow/trustgraph/query/doc_embeddings/qdrant/service.py +++ b/trustgraph-flow/trustgraph/query/doc_embeddings/qdrant/service.py @@ -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 diff --git a/trustgraph-flow/trustgraph/retrieval/document_rag/document_rag.py b/trustgraph-flow/trustgraph/retrieval/document_rag/document_rag.py index 4fc4850a..5e3c9b41 100644 --- a/trustgraph-flow/trustgraph/retrieval/document_rag/document_rag.py +++ b/trustgraph-flow/trustgraph/retrieval/document_rag/document_rag.py @@ -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) diff --git a/trustgraph-flow/trustgraph/retrieval/document_rag/rag.py b/trustgraph-flow/trustgraph/retrieval/document_rag/rag.py index fd36c7df..8c478874 100755 --- a/trustgraph-flow/trustgraph/retrieval/document_rag/rag.py +++ b/trustgraph-flow/trustgraph/retrieval/document_rag/rag.py @@ -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__) diff --git a/trustgraph-flow/trustgraph/retrieval/graph_rag/graph_rag.py b/trustgraph-flow/trustgraph/retrieval/graph_rag/graph_rag.py index 41518a61..6879023a 100644 --- a/trustgraph-flow/trustgraph/retrieval/graph_rag/graph_rag.py +++ b/trustgraph-flow/trustgraph/retrieval/graph_rag/graph_rag.py @@ -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: