mirror of
https://github.com/trustgraph-ai/trustgraph.git
synced 2026-07-21 11:11:03 +02:00
Doc RAG working
This commit is contained in:
parent
6be7b30633
commit
c9913297d2
14 changed files with 157 additions and 237 deletions
4
Makefile
4
Makefile
|
|
@ -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 \
|
||||||
|
|
|
||||||
|
|
@ -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?"""
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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',
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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"
|
||||||
|
|
|
||||||
|
|
@ -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"
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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__)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue