Updated query/store

This commit is contained in:
Cyber MacGeddon 2025-04-19 15:51:26 +01:00
parent 94f692ff93
commit 4943b87f7a
8 changed files with 162 additions and 148 deletions

View file

@ -20,4 +20,7 @@ from . prompt_client import PromptClientSpec
from . triples_store_service import TriplesStoreService from . triples_store_service import TriplesStoreService
from . graph_embeddings_store_service import GraphEmbeddingsStoreService from . graph_embeddings_store_service import GraphEmbeddingsStoreService
from . document_embeddings_store_service import DocumentEmbeddingsStoreService from . document_embeddings_store_service import DocumentEmbeddingsStoreService
from . triples_query_service import TriplesQueryService
from . graph_embeddings_query_service import GraphEmbeddingsQueryService
from . document_embeddings_query_service import DocumentEmbeddingsQueryService

View file

@ -0,0 +1,84 @@
"""
Document embeddings query service. Input is vectors. Output is list of
embeddings.
"""
from .. schema import DocumentEmbeddingsRequest, DocumentEmbeddingsResponse
from .. schema import Error, Value
from . flow_processor import FlowProcessor
from . consumer_spec import ConsumerSpec
from . producer_spec import ProducerSpec
default_ident = "ge-query"
class DocumentEmbeddingsQueryService(FlowProcessor):
def __init__(self, **params):
id = params.get("id")
super(DocumentEmbeddingsQueryService, self).__init__(
**params | { "id": id }
)
self.register_specification(
ConsumerSpec(
name = "request",
schema = DocumentEmbeddingsRequest,
handler = self.on_message
)
)
self.register_specification(
ProducerSpec(
name = "response",
schema = DocumentEmbeddingsResponse,
)
)
async def on_message(self, msg, consumer, flow):
try:
request = msg.value()
# Sender-produced ID
id = msg.properties()["id"]
print(f"Handling input {id}...", flush=True)
entities = self.query_document_embeddings(request)
print("Send response...", flush=True)
r = DocumentEmbeddingsResponse(entities=entities, error=None)
await flow("response").send(r, properties={"id": id})
print("Done.", flush=True)
except Exception as e:
print(f"Exception: {e}")
print("Send error response...", flush=True)
r = DocumentEmbeddingsResponse(
error=Error(
type = "document-embeddings-query-error",
message = str(e),
),
response=None,
)
await flow("response").send(r, properties={"id": id})
@staticmethod
def add_args(parser):
FlowProcessor.add_args(parser)
def run():
Processor.launch(default_ident, __doc__)

View file

@ -4,11 +4,12 @@ Graph embeddings query service. Input is vectors. Output is list of
embeddings. embeddings.
""" """
from .... schema import GraphEmbeddingsRequest, GraphEmbeddingsResponse from .. schema import GraphEmbeddingsRequest, GraphEmbeddingsResponse
from .... schema import Error, Value from .. schema import Error, Value
from .... base import FlowProcessor, GraphEmbeddingsQueryService, ConsumerSpec from . flow_processor import FlowProcessor
from .... base import ProducerSpec from . consumer_spec import ConsumerSpec
from . producer_spec import ProducerSpec
default_ident = "ge-query" default_ident = "ge-query"
@ -75,7 +76,7 @@ class GraphEmbeddingsQueryService(FlowProcessor):
@staticmethod @staticmethod
def add_args(parser): def add_args(parser):
GraphEmbeddingsQueryService.add_args(parser) FlowProcessor.add_args(parser)
def run(): def run():

View file

@ -4,11 +4,12 @@ Triples query service. Input is a (s, p, o) triple, some values may be
null. Output is a list of triples. null. Output is a list of triples.
""" """
from .... schema import TriplesQueryRequest, TriplesQueryResponse, Error from .. schema import TriplesQueryRequest, TriplesQueryResponse, Error
from .... schema import Value, Triple from .. schema import Value, Triple
from .... base import FlowProcessor, TriplesQueryService, ConsumerSpec from . flow_processor import FlowProcessor
from .... base import ProducerSpec from . consumer_spec import ConsumerSpec
from . producer_spec import ProducerSpec
default_ident = "triples-query" default_ident = "triples-query"

View file

@ -7,45 +7,32 @@ of chunks
from qdrant_client import QdrantClient from qdrant_client import QdrantClient
from qdrant_client.models import PointStruct from qdrant_client.models import PointStruct
from qdrant_client.models import Distance, VectorParams from qdrant_client.models import Distance, VectorParams
import uuid
from .... schema import DocumentEmbeddingsRequest, DocumentEmbeddingsResponse from .... schema import DocumentEmbeddingsResponse
from .... schema import Error, Value from .... schema import Error, Value
from .... schema import document_embeddings_request_queue from .... base import DocumentEmbeddingsQueryService
from .... schema import document_embeddings_response_queue
from .... base import ConsumerProducer
module = "de-query" default_ident = "de-query"
default_input_queue = document_embeddings_request_queue
default_output_queue = document_embeddings_response_queue
default_subscriber = module
default_store_uri = 'http://localhost:6333' default_store_uri = 'http://localhost:6333'
class Processor(ConsumerProducer): class Processor(DocumentEmbeddingsQueryService):
def __init__(self, **params): 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)
store_uri = params.get("store_uri", default_store_uri) store_uri = params.get("store_uri", default_store_uri)
#optional api key #optional api key
api_key = params.get("api_key", None) api_key = params.get("api_key", None)
super(Processor, self).__init__( super(Processor, self).__init__(
**params | { **params | {
"input_queue": input_queue,
"output_queue": output_queue,
"subscriber": subscriber,
"input_schema": DocumentEmbeddingsRequest,
"output_schema": DocumentEmbeddingsResponse,
"store_uri": store_uri, "store_uri": store_uri,
"api_key": api_key, "api_key": api_key,
} }
) )
self.client = 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 handle(self, msg):
@ -68,7 +55,7 @@ class Processor(ConsumerProducer):
str(dim) str(dim)
) )
search_result = self.client.query_points( search_result = self.qdrant.query_points(
collection_name=collection, collection_name=collection,
query=vec, query=vec,
limit=v.limit, limit=v.limit,
@ -79,37 +66,17 @@ class Processor(ConsumerProducer):
ent = r.payload["doc"] ent = r.payload["doc"]
chunks.append(ent) chunks.append(ent)
print("Send response...", flush=True) return chunks
r = DocumentEmbeddingsResponse(documents=chunks, error=None)
await self.send(r, properties={"id": id})
print("Done.", flush=True)
except Exception as e: except Exception as e:
print(f"Exception: {e}") print(f"Exception: {e}")
raise e
print("Send error response...", flush=True)
r = DocumentEmbeddingsResponse(
error=Error(
type = "llm-error",
message = str(e),
),
documents=None,
)
await self.send(r, properties={"id": id})
self.consumer.acknowledge(msg)
@staticmethod @staticmethod
def add_args(parser): def add_args(parser):
ConsumerProducer.add_args( DocumentEmbeddingsQueryService.add_args(parser)
parser, default_input_queue, default_subscriber,
default_output_queue,
)
parser.add_argument( parser.add_argument(
'-t', '--store-uri', '-t', '--store-uri',
@ -125,5 +92,5 @@ class Processor(ConsumerProducer):
def run(): def run():
Processor.launch(module, __doc__) Processor.launch(default_ident, __doc__)

View file

@ -12,6 +12,8 @@ from .... schema import GraphEmbeddingsResponse
from .... schema import Error, Value from .... schema import Error, Value
from .... base import GraphEmbeddingsQueryService from .... base import GraphEmbeddingsQueryService
default_ident = "ge-query"
default_store_uri = 'http://localhost:6333' default_store_uri = 'http://localhost:6333'
class Processor(GraphEmbeddingsQueryService): class Processor(GraphEmbeddingsQueryService):
@ -19,6 +21,8 @@ class Processor(GraphEmbeddingsQueryService):
def __init__(self, **params): def __init__(self, **params):
store_uri = params.get("store_uri", default_store_uri) store_uri = params.get("store_uri", default_store_uri)
#optional api key
api_key = params.get("api_key", None) api_key = params.get("api_key", None)
super(Processor, self).__init__( super(Processor, self).__init__(

View file

@ -7,38 +7,24 @@ null. Output is a list of triples.
from .... direct.cassandra import TrustGraph from .... direct.cassandra import TrustGraph
from .... schema import TriplesQueryRequest, TriplesQueryResponse, Error from .... schema import TriplesQueryRequest, TriplesQueryResponse, Error
from .... schema import Value, Triple from .... schema import Value, Triple
from .... schema import triples_request_queue from .... base import TriplesQueryService
from .... schema import triples_response_queue
from .... base import ConsumerProducer
module = "triples-query" default_ident = "triples-query"
default_input_queue = triples_request_queue
default_output_queue = triples_response_queue
default_subscriber = module
default_graph_host='localhost' default_graph_host='localhost'
class Processor(ConsumerProducer): class Processor(TriplesQueryService):
def __init__(self, **params): 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)
graph_host = params.get("graph_host", default_graph_host) graph_host = params.get("graph_host", default_graph_host)
graph_username = params.get("graph_username", None) graph_username = params.get("graph_username", None)
graph_password = params.get("graph_password", None) graph_password = params.get("graph_password", None)
super(Processor, self).__init__( super(Processor, self).__init__(
**params | { **params | {
"input_queue": input_queue,
"output_queue": output_queue,
"subscriber": subscriber,
"input_schema": TriplesQueryRequest,
"output_schema": TriplesQueryResponse,
"graph_host": graph_host, "graph_host": graph_host,
"graph_username": graph_username, "graph_username": graph_username,
"graph_password": graph_password,
} }
) )
@ -53,25 +39,23 @@ class Processor(ConsumerProducer):
else: else:
return Value(value=ent, is_uri=False) return Value(value=ent, is_uri=False)
async def handle(self, msg): async def query_triples(self, query):
try: try:
v = msg.value() table = (query.user, query.collection)
table = (v.user, v.collection)
if table != self.table: if table != self.table:
if self.username and self.password: if self.username and self.password:
self.tg = TrustGraph( self.tg = TrustGraph(
hosts=self.graph_host, hosts=self.graph_host,
keyspace=v.user, table=v.collection, keyspace=query.user, table=query.collection,
username=self.username, password=self.password username=self.username, password=self.password
) )
else: else:
self.tg = TrustGraph( self.tg = TrustGraph(
hosts=self.graph_host, hosts=self.graph_host,
keyspace=v.user, table=v.collection, keyspace=query.user, table=query.collection,
) )
self.table = table self.table = table
@ -82,63 +66,63 @@ class Processor(ConsumerProducer):
triples = [] triples = []
if v.s is not None: if query.s is not None:
if v.p is not None: if query.p is not None:
if v.o is not None: if query.o is not None:
resp = self.tg.get_spo( resp = self.tg.get_spo(
v.s.value, v.p.value, v.o.value, query.s.value, query.p.value, query.o.value,
limit=v.limit limit=query.limit
) )
triples.append((v.s.value, v.p.value, v.o.value)) triples.append((query.s.value, query.p.value, query.o.value))
else: else:
resp = self.tg.get_sp( resp = self.tg.get_sp(
v.s.value, v.p.value, query.s.value, query.p.value,
limit=v.limit limit=query.limit
) )
for t in resp: for t in resp:
triples.append((v.s.value, v.p.value, t.o)) triples.append((query.s.value, query.p.value, t.o))
else: else:
if v.o is not None: if query.o is not None:
resp = self.tg.get_os( resp = self.tg.get_os(
v.o.value, v.s.value, query.o.value, query.s.value,
limit=v.limit limit=query.limit
) )
for t in resp: for t in resp:
triples.append((v.s.value, t.p, v.o.value)) triples.append((query.s.value, t.p, query.o.value))
else: else:
resp = self.tg.get_s( resp = self.tg.get_s(
v.s.value, query.s.value,
limit=v.limit limit=query.limit
) )
for t in resp: for t in resp:
triples.append((v.s.value, t.p, t.o)) triples.append((query.s.value, t.p, t.o))
else: else:
if v.p is not None: if query.p is not None:
if v.o is not None: if query.o is not None:
resp = self.tg.get_po( resp = self.tg.get_po(
v.p.value, v.o.value, query.p.value, query.o.value,
limit=v.limit limit=query.limit
) )
for t in resp: for t in resp:
triples.append((t.s, v.p.value, v.o.value)) triples.append((t.s, query.p.value, query.o.value))
else: else:
resp = self.tg.get_p( resp = self.tg.get_p(
v.p.value, query.p.value,
limit=v.limit limit=query.limit
) )
for t in resp: for t in resp:
triples.append((t.s, v.p.value, t.o)) triples.append((t.s, query.p.value, t.o))
else: else:
if v.o is not None: if query.o is not None:
resp = self.tg.get_o( resp = self.tg.get_o(
v.o.value, query.o.value,
limit=v.limit limit=query.limit
) )
for t in resp: for t in resp:
triples.append((t.s, t.p, v.o.value)) triples.append((t.s, t.p, query.o.value))
else: else:
resp = self.tg.get_all( resp = self.tg.get_all(
limit=v.limit limit=query.limit
) )
for t in resp: for t in resp:
triples.append((t.s, t.p, t.o)) triples.append((t.s, t.p, t.o))
@ -152,37 +136,19 @@ class Processor(ConsumerProducer):
for t in triples for t in triples
] ]
print("Send response...", flush=True) return triples
r = TriplesQueryResponse(triples=triples, error=None)
await self.send(r, properties={"id": id})
print("Done.", flush=True) print("Done.", flush=True)
except Exception as e: except Exception as e:
print(f"Exception: {e}") print(f"Exception: {e}")
raise e
print("Send error response...", flush=True)
r = TriplesQueryResponse(
error=Error(
type = "llm-error",
message = str(e),
),
response=None,
)
await self.send(r, properties={"id": id})
self.consumer.acknowledge(msg)
@staticmethod @staticmethod
def add_args(parser): def add_args(parser):
ConsumerProducer.add_args( TriplesQueryService.add_args(parser)
parser, default_input_queue, default_subscriber,
default_output_queue,
)
parser.add_argument( parser.add_argument(
'-g', '--graph-host', '-g', '--graph-host',
@ -205,5 +171,5 @@ class Processor(ConsumerProducer):
def run(): def run():
Processor.launch(module, __doc__) Processor.launch(default_ident, __doc__)

View file

@ -8,31 +8,21 @@ from qdrant_client.models import PointStruct
from qdrant_client.models import Distance, VectorParams from qdrant_client.models import Distance, VectorParams
import uuid import uuid
from .... schema import DocumentEmbeddings from .... base import DocumentEmbeddingsStoreService
from .... schema import document_embeddings_store_queue
from .... log_level import LogLevel
from .... base import Consumer
module = "de-write" default_ident = "de-write"
default_input_queue = document_embeddings_store_queue
default_subscriber = module
default_store_uri = 'http://localhost:6333' default_store_uri = 'http://localhost:6333'
class Processor(Consumer): class Processor(DocumentEmbeddingsStoreService):
def __init__(self, **params): def __init__(self, **params):
input_queue = params.get("input_queue", default_input_queue)
subscriber = params.get("subscriber", default_subscriber)
store_uri = params.get("store_uri", default_store_uri) store_uri = params.get("store_uri", default_store_uri)
api_key = params.get("api_key", None) api_key = params.get("api_key", None)
super(Processor, self).__init__( super(Processor, self).__init__(
**params | { **params | {
"input_queue": input_queue,
"subscriber": subscriber,
"input_schema": DocumentEmbeddings,
"store_uri": store_uri, "store_uri": store_uri,
"api_key": api_key, "api_key": api_key,
} }
@ -40,13 +30,11 @@ class Processor(Consumer):
self.last_collection = None self.last_collection = None
self.client = QdrantClient(url=store_uri) self.qdrant = QdrantClient(url=store_uri, api_key=api_key)
async def handle(self, msg): async def store_document_embeddings(self, message):
v = msg.value() for emb in message.chunks:
for emb in v.chunks:
chunk = emb.chunk.decode("utf-8") chunk = emb.chunk.decode("utf-8")
if chunk == "": return if chunk == "": return
@ -55,7 +43,7 @@ class Processor(Consumer):
dim = len(vec) dim = len(vec)
collection = ( collection = (
"d_" + v.metadata.user + "_" + v.metadata.collection + "_" + "d_" + message.metadata.user + "_" + message.metadata.collection + "_" +
str(dim) str(dim)
) )
@ -92,7 +80,7 @@ class Processor(Consumer):
@staticmethod @staticmethod
def add_args(parser): def add_args(parser):
Consumer.add_args( DocumentEmbeddingsStoreService.add_args(
parser, default_input_queue, default_subscriber, parser, default_input_queue, default_subscriber,
) )
@ -110,5 +98,5 @@ class Processor(Consumer):
def run(): def run():
Processor.launch(module, __doc__) Processor.launch(default_ident, __doc__)