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 . graph_embeddings_store_service import GraphEmbeddingsStoreService
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.
"""
from .... schema import GraphEmbeddingsRequest, GraphEmbeddingsResponse
from .... schema import Error, Value
from .. schema import GraphEmbeddingsRequest, GraphEmbeddingsResponse
from .. schema import Error, Value
from .... base import FlowProcessor, GraphEmbeddingsQueryService, ConsumerSpec
from .... base import ProducerSpec
from . flow_processor import FlowProcessor
from . consumer_spec import ConsumerSpec
from . producer_spec import ProducerSpec
default_ident = "ge-query"
@ -75,7 +76,7 @@ class GraphEmbeddingsQueryService(FlowProcessor):
@staticmethod
def add_args(parser):
GraphEmbeddingsQueryService.add_args(parser)
FlowProcessor.add_args(parser)
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.
"""
from .... schema import TriplesQueryRequest, TriplesQueryResponse, Error
from .... schema import Value, Triple
from .. schema import TriplesQueryRequest, TriplesQueryResponse, Error
from .. schema import Value, Triple
from .... base import FlowProcessor, TriplesQueryService, ConsumerSpec
from .... base import ProducerSpec
from . flow_processor import FlowProcessor
from . consumer_spec import ConsumerSpec
from . producer_spec import ProducerSpec
default_ident = "triples-query"

View file

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

View file

@ -7,38 +7,24 @@ null. Output is a list of triples.
from .... direct.cassandra import TrustGraph
from .... schema import TriplesQueryRequest, TriplesQueryResponse, Error
from .... schema import Value, Triple
from .... schema import triples_request_queue
from .... schema import triples_response_queue
from .... base import ConsumerProducer
from .... base import TriplesQueryService
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'
class Processor(ConsumerProducer):
class Processor(TriplesQueryService):
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_username = params.get("graph_username", None)
graph_password = params.get("graph_password", None)
super(Processor, self).__init__(
**params | {
"input_queue": input_queue,
"output_queue": output_queue,
"subscriber": subscriber,
"input_schema": TriplesQueryRequest,
"output_schema": TriplesQueryResponse,
"graph_host": graph_host,
"graph_username": graph_username,
"graph_password": graph_password,
}
)
@ -53,25 +39,23 @@ class Processor(ConsumerProducer):
else:
return Value(value=ent, is_uri=False)
async def handle(self, msg):
async def query_triples(self, query):
try:
v = msg.value()
table = (v.user, v.collection)
table = (query.user, query.collection)
if table != self.table:
if self.username and self.password:
self.tg = TrustGraph(
hosts=self.graph_host,
keyspace=v.user, table=v.collection,
keyspace=query.user, table=query.collection,
username=self.username, password=self.password
)
else:
self.tg = TrustGraph(
hosts=self.graph_host,
keyspace=v.user, table=v.collection,
keyspace=query.user, table=query.collection,
)
self.table = table
@ -82,63 +66,63 @@ class Processor(ConsumerProducer):
triples = []
if v.s is not None:
if v.p is not None:
if v.o is not None:
if query.s is not None:
if query.p is not None:
if query.o is not None:
resp = self.tg.get_spo(
v.s.value, v.p.value, v.o.value,
limit=v.limit
query.s.value, query.p.value, query.o.value,
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:
resp = self.tg.get_sp(
v.s.value, v.p.value,
limit=v.limit
query.s.value, query.p.value,
limit=query.limit
)
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:
if v.o is not None:
if query.o is not None:
resp = self.tg.get_os(
v.o.value, v.s.value,
limit=v.limit
query.o.value, query.s.value,
limit=query.limit
)
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:
resp = self.tg.get_s(
v.s.value,
limit=v.limit
query.s.value,
limit=query.limit
)
for t in resp:
triples.append((v.s.value, t.p, t.o))
triples.append((query.s.value, t.p, t.o))
else:
if v.p is not None:
if v.o is not None:
if query.p is not None:
if query.o is not None:
resp = self.tg.get_po(
v.p.value, v.o.value,
limit=v.limit
query.p.value, query.o.value,
limit=query.limit
)
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:
resp = self.tg.get_p(
v.p.value,
limit=v.limit
query.p.value,
limit=query.limit
)
for t in resp:
triples.append((t.s, v.p.value, t.o))
triples.append((t.s, query.p.value, t.o))
else:
if v.o is not None:
if query.o is not None:
resp = self.tg.get_o(
v.o.value,
limit=v.limit
query.o.value,
limit=query.limit
)
for t in resp:
triples.append((t.s, t.p, v.o.value))
triples.append((t.s, t.p, query.o.value))
else:
resp = self.tg.get_all(
limit=v.limit
limit=query.limit
)
for t in resp:
triples.append((t.s, t.p, t.o))
@ -152,37 +136,19 @@ class Processor(ConsumerProducer):
for t in triples
]
print("Send response...", flush=True)
r = TriplesQueryResponse(triples=triples, error=None)
await self.send(r, properties={"id": id})
return triples
print("Done.", flush=True)
except Exception as e:
print(f"Exception: {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)
raise e
@staticmethod
def add_args(parser):
ConsumerProducer.add_args(
parser, default_input_queue, default_subscriber,
default_output_queue,
)
TriplesQueryService.add_args(parser)
parser.add_argument(
'-g', '--graph-host',
@ -205,5 +171,5 @@ class Processor(ConsumerProducer):
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
import uuid
from .... schema import DocumentEmbeddings
from .... schema import document_embeddings_store_queue
from .... log_level import LogLevel
from .... base import Consumer
from .... base import DocumentEmbeddingsStoreService
module = "de-write"
default_ident = "de-write"
default_input_queue = document_embeddings_store_queue
default_subscriber = module
default_store_uri = 'http://localhost:6333'
class Processor(Consumer):
class Processor(DocumentEmbeddingsStoreService):
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)
api_key = params.get("api_key", None)
super(Processor, self).__init__(
**params | {
"input_queue": input_queue,
"subscriber": subscriber,
"input_schema": DocumentEmbeddings,
"store_uri": store_uri,
"api_key": api_key,
}
@ -40,13 +30,11 @@ class Processor(Consumer):
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 v.chunks:
for emb in message.chunks:
chunk = emb.chunk.decode("utf-8")
if chunk == "": return
@ -55,7 +43,7 @@ class Processor(Consumer):
dim = len(vec)
collection = (
"d_" + v.metadata.user + "_" + v.metadata.collection + "_" +
"d_" + message.metadata.user + "_" + message.metadata.collection + "_" +
str(dim)
)
@ -92,7 +80,7 @@ class Processor(Consumer):
@staticmethod
def add_args(parser):
Consumer.add_args(
DocumentEmbeddingsStoreService.add_args(
parser, default_input_queue, default_subscriber,
)
@ -110,5 +98,5 @@ class Processor(Consumer):
def run():
Processor.launch(module, __doc__)
Processor.launch(default_ident, __doc__)