Merge branch 'release/v0.21' into pulsar-api-support

This commit is contained in:
Tyler O 2025-02-10 17:29:00 +00:00
commit 5bdb9c1919
149 changed files with 3916 additions and 1823 deletions

View file

@ -0,0 +1,6 @@
#!/usr/bin/env python3
from trustgraph.embeddings.document_embeddings import run
run()

View file

@ -0,0 +1,6 @@
#!/usr/bin/env python3
from trustgraph.embeddings.fastembed import run
run()

View file

@ -1,6 +0,0 @@
#!/usr/bin/env python3
from trustgraph.embeddings.vectorize import run
run()

View file

@ -0,0 +1,6 @@
#!/usr/bin/env python3
from trustgraph.embeddings.graph_embeddings import run
run()

View file

@ -34,58 +34,63 @@ setuptools.setup(
python_requires='>=3.8',
download_url = "https://github.com/trustgraph-ai/trustgraph/archive/refs/tags/v" + version + ".tar.gz",
install_requires=[
"trustgraph-base>=0.18,<0.19",
"urllib3",
"rdflib",
"pymilvus",
"langchain",
"langchain-core",
"langchain-text-splitters",
"langchain-community",
"requests",
"cassandra-driver",
"pulsar-client",
"pypdf",
"qdrant-client",
"tabulate",
"trustgraph-base>=0.21,<0.22",
"aiohttp",
"anthropic",
"pyyaml",
"prometheus-client",
"cassandra-driver",
"cohere",
"openai",
"neo4j",
"tiktoken",
"cryptography",
"falkordb",
"fastembed",
"google-generativeai",
"ibis",
"jsonschema",
"aiohttp",
"langchain",
"langchain-community",
"langchain-core",
"langchain-text-splitters",
"neo4j",
"ollama",
"openai",
"pinecone[grpc]",
"falkordb",
"prometheus-client",
"pulsar-client",
"pymilvus",
"pypdf",
"pyyaml",
"qdrant-client",
"rdflib",
"requests",
"tabulate",
"tiktoken",
"urllib3",
],
scripts=[
"scripts/api-gateway",
"scripts/agent-manager-react",
"scripts/api-gateway",
"scripts/chunker-recursive",
"scripts/chunker-token",
"scripts/de-query-milvus",
"scripts/de-query-qdrant",
"scripts/de-query-pinecone",
"scripts/de-query-qdrant",
"scripts/de-write-milvus",
"scripts/de-write-qdrant",
"scripts/de-write-pinecone",
"scripts/de-write-qdrant",
"scripts/document-embeddings",
"scripts/document-rag",
"scripts/embeddings-ollama",
"scripts/embeddings-vectorize",
"scripts/embeddings-fastembed",
"scripts/ge-query-milvus",
"scripts/ge-query-pinecone",
"scripts/ge-query-qdrant",
"scripts/ge-write-milvus",
"scripts/ge-write-pinecone",
"scripts/ge-write-qdrant",
"scripts/graph-embeddings",
"scripts/graph-rag",
"scripts/kg-extract-definitions",
"scripts/kg-extract-topics",
"scripts/kg-extract-relationships",
"scripts/kg-extract-topics",
"scripts/metering",
"scripts/object-extract-row",
"scripts/oe-write-milvus",
@ -103,13 +108,13 @@ setuptools.setup(
"scripts/text-completion-ollama",
"scripts/text-completion-openai",
"scripts/triples-query-cassandra",
"scripts/triples-query-neo4j",
"scripts/triples-query-memgraph",
"scripts/triples-query-falkordb",
"scripts/triples-query-memgraph",
"scripts/triples-query-neo4j",
"scripts/triples-write-cassandra",
"scripts/triples-write-neo4j",
"scripts/triples-write-memgraph",
"scripts/triples-write-falkordb",
"scripts/triples-write-memgraph",
"scripts/triples-write-neo4j",
"scripts/wikipedia-lookup",
]
)

View file

@ -14,8 +14,6 @@ from ... schema import AgentRequest, AgentResponse, AgentStep
from ... schema import agent_request_queue, agent_response_queue
from ... schema import prompt_request_queue as pr_request_queue
from ... schema import prompt_response_queue as pr_response_queue
from ... schema import text_completion_request_queue as tc_request_queue
from ... schema import text_completion_response_queue as tc_response_queue
from ... schema import graph_rag_request_queue as gr_request_queue
from ... schema import graph_rag_response_queue as gr_response_queue
from ... clients.prompt_client import PromptClient
@ -133,12 +131,6 @@ class Processor(ConsumerProducer):
prompt_response_queue = params.get(
"prompt_response_queue", pr_response_queue
)
text_completion_request_queue = params.get(
"text_completion_request_queue", tc_request_queue
)
text_completion_response_queue = params.get(
"text_completion_response_queue", tc_response_queue
)
graph_rag_request_queue = params.get(
"graph_rag_request_queue", gr_request_queue
)
@ -155,8 +147,6 @@ class Processor(ConsumerProducer):
"output_schema": AgentResponse,
"prompt_request_queue": prompt_request_queue,
"prompt_response_queue": prompt_response_queue,
"text_completion_request_queue": tc_request_queue,
"text_completion_response_queue": tc_response_queue,
"graph_rag_request_queue": gr_request_queue,
"graph_rag_response_queue": gr_response_queue,
}
@ -170,14 +160,6 @@ class Processor(ConsumerProducer):
pulsar_api_key=self.pulsar_api_key,
)
self.llm = LlmClient(
subscriber=subscriber,
input_queue=text_completion_request_queue,
output_queue=text_completion_response_queue,
pulsar_host = self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
)
self.graph_rag = GraphRagClient(
subscriber=subscriber,
input_queue=graph_rag_request_queue,
@ -343,18 +325,6 @@ class Processor(ConsumerProducer):
help=f'Prompt response queue (default: {pr_response_queue})',
)
parser.add_argument(
'--text-completion-request-queue',
default=tc_request_queue,
help=f'Text completion request queue (default: {tc_request_queue})',
)
parser.add_argument(
'--text-completion-response-queue',
default=tc_response_queue,
help=f'Text completion response queue (default: {tc_response_queue})',
)
parser.add_argument(
'--graph-rag-request-queue',
default=gr_request_queue,

View file

@ -14,6 +14,6 @@ class TextCompletionImpl:
self.context = context
def invoke(self, **arguments):
return self.context.prompt.request(
"question", { "question": arguments.get("computation") }
"question", { "question": arguments.get("question") }
)

View file

@ -6,7 +6,7 @@ class TrustGraph:
def __init__(
self, hosts=None,
keyspace="trustgraph", table="default",
keyspace="trustgraph", table="default", username=None, password=None
):
if hosts is None:
@ -14,8 +14,13 @@ class TrustGraph:
self.keyspace = keyspace
self.table = table
self.username = username
self.cluster = Cluster(hosts)
if username and password:
auth_provider = PlainTextAuthProvider(username=username, password=password)
self.cluster = Cluster(hosts, auth_provider=auth_provider)
else:
self.cluster = Cluster(hosts)
self.session = self.cluster.connect()
self.init()

View file

@ -16,6 +16,44 @@ from . schema import document_embeddings_response_queue
LABEL="http://www.w3.org/2000/01/rdf-schema#label"
DEFINITION="http://www.w3.org/2004/02/skos/core#definition"
class Query:
def __init__(self, rag, user, collection, verbose):
self.rag = rag
self.user = user
self.collection = collection
self.verbose = verbose
def get_vector(self, query):
if self.verbose:
print("Compute embeddings...", flush=True)
qembeds = self.rag.embeddings.request(query)
if self.verbose:
print("Done.", flush=True)
return qembeds
def get_docs(self, query):
vectors = self.get_vector(query)
if self.verbose:
print("Get entities...", flush=True)
docs = self.rag.de_client.request(
vectors, limit=self.rag.doc_limit
)
if self.verbose:
print("Docs:", flush=True)
for doc in docs:
print(doc, flush=True)
return docs
class DocumentRag:
def __init__(
@ -56,7 +94,7 @@ class DocumentRag:
print("Initialising...", flush=True)
# FIXME: Configurable
self.entity_limit = 20
self.doc_limit = 20
self.de_client = DocumentEmbeddingsClient(
pulsar_host=pulsar_host,
@ -85,42 +123,16 @@ class DocumentRag:
if self.verbose:
print("Initialised", flush=True)
def get_vector(self, query):
if self.verbose:
print("Compute embeddings...", flush=True)
qembeds = self.embeddings.request(query)
if self.verbose:
print("Done.", flush=True)
return qembeds
def get_docs(self, query):
vectors = self.get_vector(query)
if self.verbose:
print("Get entities...", flush=True)
docs = self.de_client.request(
vectors, self.entity_limit
)
if self.verbose:
print("Docs:", flush=True)
for doc in docs:
print(doc, flush=True)
return docs
def query(self, query):
def query(self, query, user="trustgraph", collection="default"):
if self.verbose:
print("Construct prompt...", flush=True)
docs = self.get_docs(query)
q = Query(
rag=self, user=user, collection=collection, verbose=self.verbose
)
docs = q.get_docs(query)
if self.verbose:
print("Invoke LLM...", flush=True)

View file

@ -0,0 +1,3 @@
from . embeddings import *

View file

@ -1,5 +1,5 @@
from . vectorize import run
from . embeddings import run
if __name__ == '__main__':
run()

View file

@ -1,11 +1,13 @@
"""
Vectorizer, calls the embeddings service to get embeddings for a chunk.
Input is text chunk, output is chunk and vectors.
Document embeddings, calls the embeddings service to get embeddings for a
chunk of text. Input is chunk of text plus metadata.
Output is chunk plus embedding.
"""
from ... schema import Chunk, ChunkEmbeddings
from ... schema import chunk_ingest_queue, chunk_embeddings_ingest_queue
from ... schema import Chunk, ChunkEmbeddings, DocumentEmbeddings
from ... schema import chunk_ingest_queue
from ... schema import document_embeddings_store_queue
from ... schema import embeddings_request_queue, embeddings_response_queue
from ... clients.embeddings_client import EmbeddingsClient
from ... log_level import LogLevel
@ -14,7 +16,7 @@ from ... base import ConsumerProducer
module = ".".join(__name__.split(".")[1:-1])
default_input_queue = chunk_ingest_queue
default_output_queue = chunk_embeddings_ingest_queue
default_output_queue = document_embeddings_store_queue
default_subscriber = module
class Processor(ConsumerProducer):
@ -39,7 +41,7 @@ class Processor(ConsumerProducer):
"embeddings_response_queue": emb_response_queue,
"subscriber": subscriber,
"input_schema": Chunk,
"output_schema": ChunkEmbeddings,
"output_schema": DocumentEmbeddings,
}
)
@ -51,31 +53,35 @@ class Processor(ConsumerProducer):
subscriber=module + "-emb",
)
def emit(self, metadata, chunk, vectors):
r = ChunkEmbeddings(metadata=metadata, chunk=chunk, vectors=vectors)
self.producer.send(r)
def handle(self, msg):
v = msg.value()
print(f"Indexing {v.metadata.id}...", flush=True)
chunk = v.chunk.decode("utf-8")
try:
vectors = self.embeddings.request(chunk)
vectors = self.embeddings.request(v.chunk)
self.emit(
embeds = [
ChunkEmbeddings(
chunk=v.chunk,
vectors=vectors,
)
]
r = DocumentEmbeddings(
metadata=v.metadata,
chunk=chunk.encode("utf-8"),
vectors=vectors
chunks=embeds,
)
self.producer.send(r)
except Exception as e:
print("Exception:", e, flush=True)
# Retry
raise e
print("Done.", flush=True)
@staticmethod

View file

@ -0,0 +1,3 @@
from . processor import *

View file

@ -0,0 +1,7 @@
#!/usr/bin/env python3
from . processor import run
if __name__ == '__main__':
run()

View file

@ -0,0 +1,89 @@
"""
Embeddings service, applies an embeddings model selected from HuggingFace.
Input is text, output is embeddings vector.
"""
from ... schema import EmbeddingsRequest, EmbeddingsResponse
from ... schema import embeddings_request_queue, embeddings_response_queue
from ... log_level import LogLevel
from ... base import ConsumerProducer
from fastembed import TextEmbedding
import os
module = ".".join(__name__.split(".")[1:-1])
default_input_queue = embeddings_request_queue
default_output_queue = embeddings_response_queue
default_subscriber = module
default_model="sentence-transformers/all-MiniLM-L6-v2"
class Processor(ConsumerProducer):
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)
model = params.get("model", default_model)
super(Processor, self).__init__(
**params | {
"input_queue": input_queue,
"output_queue": output_queue,
"subscriber": subscriber,
"input_schema": EmbeddingsRequest,
"output_schema": EmbeddingsResponse,
"model": model,
}
)
self.embeddings = TextEmbedding(model_name = model)
def handle(self, msg):
v = msg.value()
# Sender-produced ID
id = msg.properties()["id"]
print(f"Handling input {id}...", flush=True)
text = v.text
vecs = self.embeddings.embed([text])
vecs = [
v.tolist()
for v in vecs
]
print("Send response...", flush=True)
r = EmbeddingsResponse(
vectors=list(vecs),
error=None,
)
self.producer.send(r, properties={"id": id})
print("Done.", flush=True)
@staticmethod
def add_args(parser):
ConsumerProducer.add_args(
parser, default_input_queue, default_subscriber,
default_output_queue,
)
parser.add_argument(
'-m', '--model',
default=default_model,
help=f'Embeddings model (default: {default_model})'
)
def run():
Processor.start(module, __doc__)

View file

@ -0,0 +1,3 @@
from . embeddings import *

View file

@ -0,0 +1,6 @@
from . embeddings import run
if __name__ == '__main__':
run()

View file

@ -0,0 +1,113 @@
"""
Graph embeddings, calls the embeddings service to get embeddings for a
set of entity contexts. Input is entity plus textual context.
Output is entity plus embedding.
"""
from ... schema import EntityContexts, EntityEmbeddings, GraphEmbeddings
from ... schema import entity_contexts_ingest_queue
from ... schema import graph_embeddings_store_queue
from ... schema import embeddings_request_queue, embeddings_response_queue
from ... clients.embeddings_client import EmbeddingsClient
from ... log_level import LogLevel
from ... base import ConsumerProducer
module = ".".join(__name__.split(".")[1:-1])
default_input_queue = entity_contexts_ingest_queue
default_output_queue = graph_embeddings_store_queue
default_subscriber = module
class Processor(ConsumerProducer):
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)
emb_request_queue = params.get(
"embeddings_request_queue", embeddings_request_queue
)
emb_response_queue = params.get(
"embeddings_response_queue", embeddings_response_queue
)
super(Processor, self).__init__(
**params | {
"input_queue": input_queue,
"output_queue": output_queue,
"embeddings_request_queue": emb_request_queue,
"embeddings_response_queue": emb_response_queue,
"subscriber": subscriber,
"input_schema": EntityContexts,
"output_schema": GraphEmbeddings,
}
)
self.embeddings = EmbeddingsClient(
pulsar_host=self.pulsar_host,
input_queue=emb_request_queue,
output_queue=emb_response_queue,
subscriber=module + "-emb",
)
def handle(self, msg):
v = msg.value()
print(f"Indexing {v.metadata.id}...", flush=True)
entities = []
try:
for entity in v.entities:
vectors = self.embeddings.request(entity.context)
entities.append(
EntityEmbeddings(
entity=entity.entity,
vectors=vectors
)
)
r = GraphEmbeddings(
metadata=v.metadata,
entities=entities,
)
self.producer.send(r)
except Exception as e:
print("Exception:", e, flush=True)
# Retry
raise e
print("Done.", flush=True)
@staticmethod
def add_args(parser):
ConsumerProducer.add_args(
parser, default_input_queue, default_subscriber,
default_output_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 request queue (default: {embeddings_response_queue})',
)
def run():
Processor.start(module, __doc__)

View file

@ -1,14 +1,15 @@
"""
Embeddings service, applies an embeddings model selected from HuggingFace.
Embeddings service, applies an embeddings model hosted on a local Ollama.
Input is text, output is embeddings vector.
"""
from langchain_community.embeddings import OllamaEmbeddings
from ... schema import EmbeddingsRequest, EmbeddingsResponse
from ... schema import embeddings_request_queue, embeddings_response_queue
from ... log_level import LogLevel
from ... base import ConsumerProducer
from ollama import Client
import os
module = ".".join(__name__.split(".")[1:-1])
@ -16,7 +17,7 @@ default_input_queue = embeddings_request_queue
default_output_queue = embeddings_response_queue
default_subscriber = module
default_model="mxbai-embed-large"
default_ollama = 'http://localhost:11434'
default_ollama = os.getenv("OLLAMA_HOST", 'http://localhost:11434')
class Processor(ConsumerProducer):
@ -26,6 +27,9 @@ class Processor(ConsumerProducer):
output_queue = params.get("output_queue", default_output_queue)
subscriber = params.get("subscriber", default_subscriber)
ollama = params.get("ollama", default_ollama)
model = params.get("model", default_model)
super(Processor, self).__init__(
**params | {
"input_queue": input_queue,
@ -33,10 +37,13 @@ class Processor(ConsumerProducer):
"subscriber": subscriber,
"input_schema": EmbeddingsRequest,
"output_schema": EmbeddingsResponse,
"ollama": ollama,
"model": model,
}
)
self.embeddings = OllamaEmbeddings(base_url=ollama, model=model)
self.client = Client(host=ollama)
self.model = model
def handle(self, msg):
@ -49,10 +56,16 @@ class Processor(ConsumerProducer):
print(f"Handling input {id}...", flush=True)
text = v.text
embeds = self.embeddings.embed_query([text])
embeds = self.client.embed(
model = self.model,
input = text
)
print("Send response...", flush=True)
r = EmbeddingsResponse(vectors=[embeds])
r = EmbeddingsResponse(
vectors=embeds.embeddings,
error=None,
)
self.producer.send(r, properties={"id": id})

View file

@ -1,3 +0,0 @@
from . vectorize import *

View file

@ -1,14 +1,17 @@
"""
Simple decoder, accepts embeddings+text chunks input, applies entity analysis to
get entity definitions which are output as graph edges.
Simple decoder, accepts text chunks input, applies entity analysis to
get entity definitions which are output as graph edges along with
entity/context definitions for embedding.
"""
import urllib.parse
import json
from pulsar.schema import JsonSchema
from .... schema import ChunkEmbeddings, Triple, Triples, Metadata, Value
from .... schema import chunk_embeddings_ingest_queue, triples_store_queue
from .... schema import Chunk, Triple, Triples, Metadata, Value
from .... schema import EntityContext, EntityContexts
from .... schema import chunk_ingest_queue, triples_store_queue
from .... schema import entity_contexts_ingest_queue
from .... schema import prompt_request_queue
from .... schema import prompt_response_queue
from .... log_level import LogLevel
@ -22,8 +25,9 @@ SUBJECT_OF_VALUE = Value(value=SUBJECT_OF, is_uri=True)
module = ".".join(__name__.split(".")[1:-1])
default_input_queue = chunk_embeddings_ingest_queue
default_input_queue = chunk_ingest_queue
default_output_queue = triples_store_queue
default_entity_context_queue = entity_contexts_ingest_queue
default_subscriber = module
class Processor(ConsumerProducer):
@ -32,6 +36,10 @@ class Processor(ConsumerProducer):
input_queue = params.get("input_queue", default_input_queue)
output_queue = params.get("output_queue", default_output_queue)
ec_queue = params.get(
"entity_context_queue",
default_entity_context_queue
)
subscriber = params.get("subscriber", default_subscriber)
pr_request_queue = params.get(
"prompt_request_queue", prompt_request_queue
@ -45,13 +53,30 @@ class Processor(ConsumerProducer):
"input_queue": input_queue,
"output_queue": output_queue,
"subscriber": subscriber,
"input_schema": ChunkEmbeddings,
"input_schema": Chunk,
"output_schema": Triples,
"prompt_request_queue": pr_request_queue,
"prompt_response_queue": pr_response_queue,
}
)
self.ec_prod = self.client.create_producer(
topic=ec_queue,
schema=JsonSchema(EntityContexts),
)
__class__.pubsub_metric.info({
"input_queue": input_queue,
"output_queue": output_queue,
"entity_context_queue": ec_queue,
"prompt_request_queue": pr_request_queue,
"prompt_response_queue": pr_response_queue,
"subscriber": subscriber,
"input_schema": Chunk.__name__,
"output_schema": Triples.__name__,
"vector_schema": EntityContexts.__name__,
})
self.prompt = PromptClient(
pulsar_host=self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
@ -80,6 +105,14 @@ class Processor(ConsumerProducer):
)
self.producer.send(t)
def emit_ecs(self, metadata, entities):
t = EntityContexts(
metadata=metadata,
entities=entities,
)
self.ec_prod.send(t)
def handle(self, msg):
v = msg.value()
@ -92,6 +125,7 @@ class Processor(ConsumerProducer):
defs = self.get_definitions(chunk)
triples = []
entities = []
# FIXME: Putting metadata into triples store is duplicated in
# relationships extractor too
@ -130,6 +164,14 @@ class Processor(ConsumerProducer):
o=Value(value=v.metadata.id, is_uri=True)
))
ec = EntityContext(
entity=s_value,
context=defn.definition,
)
entities.append(ec)
self.emit_edges(
Metadata(
id=v.metadata.id,
@ -140,6 +182,16 @@ class Processor(ConsumerProducer):
triples
)
self.emit_ecs(
Metadata(
id=v.metadata.id,
metadata=[],
user=v.metadata.user,
collection=v.metadata.collection,
),
entities
)
except Exception as e:
print("Exception: ", e, flush=True)
@ -153,6 +205,12 @@ class Processor(ConsumerProducer):
default_output_queue,
)
parser.add_argument(
'-e', '--entity-context-queue',
default=default_entity_context_queue,
help=f'Entity context queue (default: {default_entity_context_queue})'
)
parser.add_argument(
'--prompt-request-queue',
default=prompt_request_queue,

View file

@ -1,18 +1,15 @@
"""
Simple decoder, accepts vector+text chunks input, applies entity
Simple decoder, accepts text chunks input, applies entity
relationship analysis to get entity relationship edges which are output as
graph edges.
"""
import urllib.parse
import os
from pulsar.schema import JsonSchema
from .... schema import ChunkEmbeddings, Triple, Triples, GraphEmbeddings
from .... schema import Chunk, Triple, Triples
from .... schema import Metadata, Value
from .... schema import chunk_embeddings_ingest_queue, triples_store_queue
from .... schema import graph_embeddings_store_queue
from .... schema import chunk_ingest_queue, triples_store_queue
from .... schema import prompt_request_queue
from .... schema import prompt_response_queue
from .... log_level import LogLevel
@ -25,9 +22,8 @@ SUBJECT_OF_VALUE = Value(value=SUBJECT_OF, is_uri=True)
module = ".".join(__name__.split(".")[1:-1])
default_input_queue = chunk_embeddings_ingest_queue
default_input_queue = chunk_ingest_queue
default_output_queue = triples_store_queue
default_vector_queue = graph_embeddings_store_queue
default_subscriber = module
class Processor(ConsumerProducer):
@ -36,7 +32,6 @@ class Processor(ConsumerProducer):
input_queue = params.get("input_queue", default_input_queue)
output_queue = params.get("output_queue", default_output_queue)
vector_queue = params.get("vector_queue", default_vector_queue)
subscriber = params.get("subscriber", default_subscriber)
pr_request_queue = params.get(
"prompt_request_queue", prompt_request_queue
@ -50,30 +45,13 @@ class Processor(ConsumerProducer):
"input_queue": input_queue,
"output_queue": output_queue,
"subscriber": subscriber,
"input_schema": ChunkEmbeddings,
"input_schema": Chunk,
"output_schema": Triples,
"prompt_request_queue": pr_request_queue,
"prompt_response_queue": pr_response_queue,
}
)
self.vec_prod = self.client.create_producer(
topic=vector_queue,
schema=JsonSchema(GraphEmbeddings),
)
__class__.pubsub_metric.info({
"input_queue": input_queue,
"output_queue": output_queue,
"vector_queue": vector_queue,
"prompt_request_queue": pr_request_queue,
"prompt_response_queue": pr_response_queue,
"subscriber": subscriber,
"input_schema": ChunkEmbeddings.__name__,
"output_schema": Triples.__name__,
"vector_schema": GraphEmbeddings.__name__,
})
self.prompt = PromptClient(
pulsar_host=self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
@ -102,11 +80,6 @@ class Processor(ConsumerProducer):
)
self.producer.send(t)
def emit_vec(self, metadata, ent, vec):
r = GraphEmbeddings(metadata=metadata, entity=ent, vectors=vec)
self.vec_prod.send(r)
def handle(self, msg):
v = msg.value()
@ -194,12 +167,6 @@ class Processor(ConsumerProducer):
o=Value(value=v.metadata.id, is_uri=True)
))
self.emit_vec(v.metadata, s_value, v.vectors)
self.emit_vec(v.metadata, p_value, v.vectors)
if rel.o_entity:
self.emit_vec(v.metadata, o_value, v.vectors)
self.emit_edges(
Metadata(
id=v.metadata.id,
@ -223,12 +190,6 @@ class Processor(ConsumerProducer):
default_output_queue,
)
parser.add_argument(
'-c', '--vector-queue',
default=default_vector_queue,
help=f'Vector output queue (default: {default_vector_queue})'
)
parser.add_argument(
'--prompt-request-queue',
default=prompt_request_queue,

View file

@ -1,14 +1,14 @@
"""
Simple decoder, accepts embeddings+text chunks input, applies entity analysis to
get entity definitions which are output as graph edges.
Simple decoder, accepts text chunks input, applies entity analysis to
get topics which are output as graph edges.
"""
import urllib.parse
import json
from .... schema import ChunkEmbeddings, Triple, Triples, Metadata, Value
from .... schema import chunk_embeddings_ingest_queue, triples_store_queue
from .... schema import Chunk, Triple, Triples, Metadata, Value
from .... schema import chunk_ingest_queue, triples_store_queue
from .... schema import prompt_request_queue
from .... schema import prompt_response_queue
from .... log_level import LogLevel
@ -20,7 +20,7 @@ DEFINITION_VALUE = Value(value=DEFINITION, is_uri=True)
module = ".".join(__name__.split(".")[1:-1])
default_input_queue = chunk_embeddings_ingest_queue
default_input_queue = chunk_ingest_queue
default_output_queue = triples_store_queue
default_subscriber = module
@ -43,7 +43,7 @@ class Processor(ConsumerProducer):
"input_queue": input_queue,
"output_queue": output_queue,
"subscriber": subscriber,
"input_schema": ChunkEmbeddings,
"input_schema": Chunk,
"output_schema": Triples,
"prompt_request_queue": pr_request_queue,
"prompt_response_queue": pr_response_queue,

View file

@ -0,0 +1,64 @@
import asyncio
from pulsar.schema import JsonSchema
import uuid
from aiohttp import WSMsgType
from .. schema import Metadata
from .. schema import DocumentEmbeddings, ChunkEmbeddings
from .. schema import document_embeddings_store_queue
from . publisher import Publisher
from . socket import SocketEndpoint
from . serialize import to_subgraph
class DocumentEmbeddingsLoadEndpoint(SocketEndpoint):
def __init__(
self, pulsar_host, auth, path="/api/v1/load/document-embeddings",
):
super(DocumentEmbeddingsLoadEndpoint, self).__init__(
endpoint_path=path, auth=auth,
)
self.pulsar_host=pulsar_host
self.publisher = Publisher(
self.pulsar_host, document_embeddings_store_queue,
schema=JsonSchema(DocumentEmbeddings)
)
async def start(self):
self.publisher.start()
async def listener(self, ws, running):
async for msg in ws:
# On error, finish
if msg.type == WSMsgType.ERROR:
break
else:
data = msg.json()
elt = DocumentEmbeddings(
metadata=Metadata(
id=data["metadata"]["id"],
metadata=to_subgraph(data["metadata"]["metadata"]),
user=data["metadata"]["user"],
collection=data["metadata"]["collection"],
),
chunks=[
ChunkEmbeddings(
chunk=de["chunk"].encode("utf-8"),
vectors=de["vectors"],
)
for de in data["chunks"]
],
)
self.publisher.send(None, elt)
running.stop()

View file

@ -0,0 +1,72 @@
import asyncio
import queue
from pulsar.schema import JsonSchema
import uuid
from .. schema import DocumentEmbeddings
from .. schema import document_embeddings_store_queue
from . subscriber import Subscriber
from . socket import SocketEndpoint
from . serialize import serialize_document_embeddings
class DocumentEmbeddingsStreamEndpoint(SocketEndpoint):
def __init__(
self, pulsar_host, auth, path="/api/v1/stream/document-embeddings"
):
super(DocumentEmbeddingsStreamEndpoint, self).__init__(
endpoint_path=path, auth=auth,
)
self.pulsar_host=pulsar_host
self.subscriber = Subscriber(
self.pulsar_host, document_embeddings_store_queue,
"api-gateway", "api-gateway",
schema=JsonSchema(DocumentEmbeddings)
)
async def listener(self, ws, running):
worker = asyncio.create_task(
self.async_thread(ws, running)
)
await super(DocumentEmbeddingsStreamEndpoint, self).listener(
ws, running
)
await worker
async def start(self):
self.subscriber.start()
async def async_thread(self, ws, running):
id = str(uuid.uuid4())
q = self.subscriber.subscribe_all(id)
while running.get():
try:
resp = await asyncio.to_thread(q.get, timeout=0.5)
await ws.send_json(serialize_document_embeddings(resp))
except TimeoutError:
continue
except queue.Empty:
continue
except Exception as e:
print(f"Exception: {str(e)}", flush=True)
break
self.subscriber.unsubscribe_all(id)
running.stop()

View file

@ -1,7 +1,7 @@
import base64
from .. schema import Document
from .. schema import Document, Metadata
from .. schema import document_ingest_queue
from . sender import ServiceSender
@ -19,25 +19,24 @@ class DocumentLoadSender(ServiceSender):
def to_request(self, body):
if "metadata" in data:
metadata = to_subgraph(data["metadata"])
if "metadata" in body:
metadata = to_subgraph(body["metadata"])
else:
metadata = []
# Doing a base64 decoe/encode here to make sure the
# content is valid base64
doc = base64.b64decode(data["data"])
doc = base64.b64decode(body["data"])
print("Document received")
return Document(
metadata=Metadata(
id=data.get("id"),
id=body.get("id"),
metadata=metadata,
user=data.get("user", "trustgraph"),
collection=data.get("collection", "default"),
user=body.get("user", "trustgraph"),
collection=body.get("collection", "default"),
),
data=base64.b64encode(doc).decode("utf-8")
)

View file

@ -0,0 +1,30 @@
from .. schema import DocumentRagQuery, DocumentRagResponse
from .. schema import document_rag_request_queue
from .. schema import document_rag_response_queue
from . endpoint import ServiceEndpoint
from . requestor import ServiceRequestor
class DocumentRagRequestor(ServiceRequestor):
def __init__(self, pulsar_host, timeout, auth):
super(DocumentRagRequestor, self).__init__(
pulsar_host=pulsar_host,
request_queue=document_rag_request_queue,
response_queue=document_rag_response_queue,
request_schema=DocumentRagQuery,
response_schema=DocumentRagResponse,
timeout=timeout,
)
def to_request(self, body):
return DocumentRagQuery(
query=body["query"],
user=body.get("user", "trustgraph"),
collection=body.get("collection", "default"),
)
def from_response(self, message):
return { "response": message.response }, True

View file

@ -5,7 +5,7 @@ import uuid
from aiohttp import WSMsgType
from .. schema import Metadata
from .. schema import GraphEmbeddings
from .. schema import GraphEmbeddings, EntityEmbeddings
from .. schema import graph_embeddings_store_queue
from . publisher import Publisher
@ -52,8 +52,13 @@ class GraphEmbeddingsLoadEndpoint(SocketEndpoint):
user=data["metadata"]["user"],
collection=data["metadata"]["collection"],
),
entity=to_value(data["entity"]),
vectors=data["vectors"],
entities=[
EntityEmbeddings(
entity=to_value(ent["entity"]),
vectors=ent["vectors"],
)
for ent in data["entities"]
]
)
self.publisher.send(None, elt)

View file

@ -31,6 +31,16 @@ class GraphEmbeddingsStreamEndpoint(SocketEndpoint):
schema=JsonSchema(GraphEmbeddings)
)
async def listener(self, ws, running):
worker = asyncio.create_task(
self.async_thread(ws, running)
)
await super(GraphEmbeddingsStreamEndpoint, self).listener(ws, running)
await worker
async def start(self):
self.subscriber.start()
@ -46,6 +56,9 @@ class GraphEmbeddingsStreamEndpoint(SocketEndpoint):
resp = await asyncio.to_thread(q.get, timeout=0.5)
await ws.send_json(serialize_graph_embeddings(resp))
except TimeoutError:
continue
except queue.Empty:
continue

View file

@ -0,0 +1,73 @@
#
# This provides a Prometheus endpoint on the api-gateway. It proxies
# HTTP GET requests to Prometheus.
#
import aiohttp
from aiohttp import web
import asyncio
from pulsar.schema import JsonSchema
import uuid
import logging
logger = logging.getLogger("endpoint")
logger.setLevel(logging.INFO)
class MetricsEndpoint:
def __init__(self, prometheus_url, endpoint_path, auth):
self.prometheus_url = prometheus_url
self.path = endpoint_path
self.auth = auth
self.operation = "service"
async def start(self):
pass
def add_routes(self, app):
app.add_routes([
web.get(self.path + "/{path:.*}", self.handle),
])
async def handle(self, request):
print(request.path, "...")
try:
ht = request.headers["Authorization"]
tokens = ht.split(" ", 2)
if tokens[0] != "Bearer":
return web.HTTPUnauthorized()
token = tokens[1]
except:
token = ""
if not self.auth.permitted(token, self.operation):
return web.HTTPUnauthorized()
try:
path = request.match_info["path"]
async with aiohttp.ClientSession() as session:
url = (
self.prometheus_url + "/api/v1/" + path + "?" +
request.query_string
)
async with session.get(url) as resp:
return web.Response(
status=resp.status,
text=await resp.text()
)
except Exception as e:
logging.error(f"Exception: {e}")
raise web.HTTPInternalServerError()

View file

@ -7,13 +7,14 @@ import threading
class Publisher:
def __init__(self, pulsar_host, topic, schema=None, max_size=10,
chunking_enabled=False, pulsar_api_key=None):
chunking_enabled=True, listener=None, pulsar_api_key=None):
self.pulsar_host = pulsar_host
self.pulsar_api_key = pulsar_api_key,
self.topic = topic
self.schema = schema
self.q = queue.Queue(maxsize=max_size)
self.chunking_enabled = chunking_enabled
self.listener_name = listener
def start(self):
self.task = threading.Thread(target=self.run)
@ -28,11 +29,13 @@ class Publisher:
if self.pulsar_api_key:
client = pulsar.Client(
self.pulsar_host,
listener_name=self.listener_name,
authentication=pulsar.AuthenticationToken(self.pulsar_api_key)
)
else:
client = pulsar.Client(
self.pulsar_host,
listener_name=self.listener_name
)
producer = client.create_producer(

View file

@ -63,12 +63,18 @@ class ServiceRequestor:
while True:
try:
resp = await asyncio.to_thread(q.get, timeout=self.timeout)
resp = await asyncio.to_thread(
q.get,
timeout=self.timeout
)
except Exception as e:
raise RuntimeError("Timeout")
if resp.error:
return { "error": resp.error.message }
err = { "error": resp.error.message }
if responder:
await responder(err, True)
return err
resp, fin = self.from_response(resp)
@ -84,7 +90,10 @@ class ServiceRequestor:
logging.error(f"Exception: {e}")
return { "error": str(e) }
err = { "error": str(e) }
if responder:
await responder(err, True)
return err
finally:
self.sub.unsubscribe(id)

View file

@ -48,5 +48,11 @@ class ServiceSender:
logging.error(f"Exception: {e}")
return { "error": str(e) }
err = { "error": str(e) }
if responder:
await responder(err, True)
return err

View file

@ -51,7 +51,29 @@ def serialize_graph_embeddings(message):
"user": message.metadata.user,
"collection": message.metadata.collection,
},
"vectors": message.vectors,
"entity": serialize_value(message.entity),
"entities": [
{
"vectors": entity.vectors,
"entity": serialize_value(entity.entity),
}
for entity in message.entities
],
}
def serialize_document_embeddings(message):
return {
"metadata": {
"id": message.metadata.id,
"metadata": serialize_subgraph(message.metadata.metadata),
"user": message.metadata.user,
"collection": message.metadata.collection,
},
"chunks": [
{
"vectors": chunk.vectors,
"chunk": chunk.chunk.decode("utf-8"),
}
for chunk in message.chunks
],
}

View file

@ -31,6 +31,7 @@ from . subscriber import Subscriber
from . text_completion import TextCompletionRequestor
from . prompt import PromptRequestor
from . graph_rag import GraphRagRequestor
from . document_rag import DocumentRagRequestor
from . triples_query import TriplesQueryRequestor
from . graph_embeddings_query import GraphEmbeddingsQueryRequestor
from . embeddings import EmbeddingsRequestor
@ -40,11 +41,14 @@ from . dbpedia import DbpediaRequestor
from . internet_search import InternetSearchRequestor
from . triples_stream import TriplesStreamEndpoint
from . graph_embeddings_stream import GraphEmbeddingsStreamEndpoint
from . document_embeddings_stream import DocumentEmbeddingsStreamEndpoint
from . triples_load import TriplesLoadEndpoint
from . graph_embeddings_load import GraphEmbeddingsLoadEndpoint
from . document_embeddings_load import DocumentEmbeddingsLoadEndpoint
from . mux import MuxEndpoint
from . document_load import DocumentLoadSender
from . text_load import TextLoadSender
from . metrics import MetricsEndpoint
from . endpoint import ServiceEndpoint
from . auth import Authenticator
@ -54,6 +58,7 @@ logger.setLevel(logging.INFO)
default_pulsar_host = os.getenv("PULSAR_HOST", "pulsar://pulsar:6650")
default_pulsar_api_key = os.getenv("PULSAR_API_KEY", None)
default_prometheus_url = os.getenv("PROMETHEUS_URL", "http://prometheus:9090")
default_timeout = 600
default_port = 8088
default_api_token = os.getenv("GATEWAY_SECRET", "")
@ -72,6 +77,13 @@ class Api:
self.pulsar_host = config.get("pulsar_host", default_pulsar_host)
self.pulsar_api_key = config.get("pulsar_api_key", default_pulsar_api_key)
self.prometheus_url = config.get(
"prometheus_url", default_prometheus_url,
)
if not self.prometheus_url.endswith("/"):
self.prometheus_url += "/"
api_token = config.get("api_token", default_api_token)
# Token not set, or token equal empty string means no auth
@ -93,6 +105,10 @@ class Api:
pulsar_host=self.pulsar_host, timeout=self.timeout,
auth = self.auth, pulsar_api_key=self.pulsar_api_key,
),
"document-rag": DocumentRagRequestor(
pulsar_host=self.pulsar_host, timeout=self.timeout,
auth = self.auth,
),
"triples-query": TriplesQueryRequestor(
pulsar_host=self.pulsar_host, timeout=self.timeout,
auth = self.auth, pulsar_api_key=self.pulsar_api_key,
@ -142,6 +158,10 @@ class Api:
endpoint_path = "/api/v1/graph-rag", auth=self.auth,
requestor = self.services["graph-rag"],
),
ServiceEndpoint(
endpoint_path = "/api/v1/document-rag", auth=self.auth,
requestor = self.services["document-rag"],
),
ServiceEndpoint(
endpoint_path = "/api/v1/triples-query", auth=self.auth,
requestor = self.services["triples-query"],
@ -189,6 +209,10 @@ class Api:
pulsar_api_key=self.pulsar_api_key,
auth = self.auth,
),
DocumentEmbeddingsStreamEndpoint(
pulsar_host=self.pulsar_host,
auth = self.auth,
),
TriplesLoadEndpoint(
pulsar_host=self.pulsar_host,
auth = self.auth,
@ -199,12 +223,21 @@ class Api:
pulsar_api_key=self.pulsar_api_key,
auth = self.auth,
),
DocumentEmbeddingsLoadEndpoint(
pulsar_host=self.pulsar_host,
auth = self.auth,
),
MuxEndpoint(
pulsar_host=self.pulsar_host,
auth = self.auth,
services = self.services,
pulsar_api_key=self.pulsar_api_key,
),
MetricsEndpoint(
endpoint_path = "/api/v1/metrics",
prometheus_url = self.prometheus_url,
auth = self.auth,
),
]
for ep in self.endpoints:
@ -239,6 +272,12 @@ def run():
help=f'Pulsar API key',
)
parser.add_argument(
'-m', '--prometheus-url',
default=default_prometheus_url,
help=f'Prometheus URL (default: {default_prometheus_url})',
)
parser.add_argument(
'--port',
type=int,

View file

@ -44,7 +44,10 @@ class SocketEndpoint:
return web.HTTPUnauthorized()
running = Running()
ws = web.WebSocketResponse()
# 50MB max message size
ws = web.WebSocketResponse(max_msg_size=52428800)
await ws.prepare(request)
try:

View file

@ -7,7 +7,7 @@ import time
class Subscriber:
def __init__(self, pulsar_host, topic, subscription, consumer_name, pulsar_api_key=None,
schema=None, max_size=100):
schema=None, max_size=100, listener=None):
self.pulsar_host = pulsar_host
self.pulsar_api_key = pulsar_api_key
self.topic = topic
@ -18,6 +18,7 @@ class Subscriber:
self.full = {}
self.max_size = max_size
self.lock = threading.Lock()
self.listener_name = listener
def start(self):
self.task = threading.Thread(target=self.run)
@ -32,12 +33,14 @@ class Subscriber:
if self.pulsar_api_key:
auth = pulsar.AuthenticationToken(self.pulsar_api_key)
client = pulsar.Client(
self.pulsar_host,
authentication=auth,
self.pulsar_host,
authentication=auth,
listener_name=self.listener_name,
)
else:
client = pulsar.Client(
self.pulsar_host,
self.pulsar_host,
listener_name=self.listener_name,
)
consumer = client.subscribe(

View file

@ -37,7 +37,7 @@ class TextLoadSender(ServiceSender):
return TextDocument(
metadata=Metadata(
id=body.get("id"),
metabody=metadata,
metadata=metadata,
user=body.get("user", "trustgraph"),
collection=body.get("collection", "default"),
),

View file

@ -29,6 +29,16 @@ class TriplesStreamEndpoint(SocketEndpoint):
schema=JsonSchema(Triples)
)
async def listener(self, ws, running):
worker = asyncio.create_task(
self.async_thread(ws, running)
)
await super(TriplesStreamEndpoint, self).listener(ws, running)
await worker
async def start(self):
self.subscriber.start()
@ -44,6 +54,9 @@ class TriplesStreamEndpoint(SocketEndpoint):
resp = await asyncio.to_thread(q.get, timeout=0.5)
await ws.send_json(serialize_triples(resp))
except TimeoutError:
continue
except queue.Empty:
continue

View file

@ -158,25 +158,15 @@ class Processor(ConsumerProducer):
except TooManyRequests:
print("Send rate limit response...", flush=True)
print("Rate limit...")
r = TextCompletionResponse(
error=Error(
type = "rate-limit",
message = str(e),
),
response=None,
in_token=None,
out_token=None,
model=None,
)
self.producer.send(r, properties={"id": id})
self.consumer.acknowledge(msg)
# Leave rate limit retries to the base handler
raise TooManyRequests()
except Exception as e:
# Apart from rate limits, treat all exceptions as unrecoverable
print(f"Exception: {e}")
print("Send error response...", flush=True)

View file

@ -4,10 +4,9 @@ Simple LLM service, performs text prompt completion using the Azure
OpenAI endpoit service. Input is prompt, output is response.
"""
import requests
import json
from prometheus_client import Histogram
from openai import AzureOpenAI
from openai import AzureOpenAI, RateLimitError
import os
from .... schema import TextCompletionRequest, TextCompletionResponse, Error
@ -126,30 +125,27 @@ class Processor(ConsumerProducer):
print(f"Output Tokens: {outputtokens}", flush=True)
print("Send response...", flush=True)
r = TextCompletionResponse(response=resp.choices[0].message.content, error=None, in_token=inputtokens, out_token=outputtokens, model=self.model)
self.producer.send(r, properties={"id": id})
except TooManyRequests:
print("Send rate limit response...", flush=True)
r = TextCompletionResponse(
error=Error(
type = "rate-limit",
message = str(e),
),
response=None,
in_token=None,
out_token=None,
model=None,
response=resp.choices[0].message.content,
error=None,
in_token=inputtokens,
out_token=outputtokens,
model=self.model
)
self.producer.send(r, properties={"id": id})
self.consumer.acknowledge(msg)
except RateLimitError:
print("Send rate limit response...", flush=True)
# Leave rate limit retries to the base handler
raise TooManyRequests()
except Exception as e:
# Apart from rate limits, treat all exceptions as unrecoverable
print(f"Exception: {e}")
print("Send error response...", flush=True)

View file

@ -87,8 +87,6 @@ class Processor(ConsumerProducer):
try:
# FIXME: Rate limits?
with __class__.text_completion_metric.time():
response = message = self.claude.messages.create(
@ -117,34 +115,26 @@ class Processor(ConsumerProducer):
print(f"Output Tokens: {outputtokens}", flush=True)
print("Send response...", flush=True)
r = TextCompletionResponse(response=resp, error=None, in_token=inputtokens, out_token=outputtokens, model=self.model)
r = TextCompletionResponse(
response=resp,
error=None,
in_token=inputtokens,
out_token=outputtokens,
model=self.model
)
self.send(r, properties={"id": id})
print("Done.", flush=True)
# FIXME: Wrong exception, don't know what this LLM throws
# for a rate limit
except TooManyRequests:
except anthropic.RateLimitError:
print("Send rate limit response...", flush=True)
r = TextCompletionResponse(
error=Error(
type = "rate-limit",
message = str(e),
),
response=None,
in_token=None,
out_token=None,
model=None,
)
self.producer.send(r, properties={"id": id})
self.consumer.acknowledge(msg)
# Leave rate limit retries to the base handler
raise TooManyRequests()
except Exception as e:
# Apart from rate limits, treat all exceptions as unrecoverable
print(f"Exception: {e}")
print("Send error response...", flush=True)

View file

@ -112,27 +112,15 @@ class Processor(ConsumerProducer):
# FIXME: Wrong exception, don't know what this LLM throws
# for a rate limit
except TooManyRequests:
except cohere.TooManyRequestsError:
print("Send rate limit response...", flush=True)
r = TextCompletionResponse(
error=Error(
type = "rate-limit",
message = str(e),
),
response=None,
in_token=None,
out_token=None,
model=None,
)
self.producer.send(r, properties={"id": id})
self.consumer.acknowledge(msg)
# Leave rate limit retries to the base handler
raise TooManyRequests()
except Exception as e:
# Apart from rate limits, treat all exceptions as unrecoverable
print(f"Exception: {e}")
print("Send error response...", flush=True)

View file

@ -88,7 +88,8 @@ class Processor(ConsumerProducer):
HarmCategory.HARM_CATEGORY_HARASSMENT: block_level,
HarmCategory.HARM_CATEGORY_SEXUALLY_EXPLICIT: block_level,
HarmCategory.HARM_CATEGORY_DANGEROUS_CONTENT: block_level,
# There is a documentation conflict on whether or not CIVIC_INTEGRITY is a valid category
# There is a documentation conflict on whether or not
# CIVIC_INTEGRITY is a valid category
# HarmCategory.HARM_CATEGORY_CIVIC_INTEGRITY: block_level,
}
@ -122,8 +123,6 @@ class Processor(ConsumerProducer):
try:
# FIXME: Rate limits?
with __class__.text_completion_metric.time():
chat_session = self.llm.start_chat(
@ -140,35 +139,30 @@ class Processor(ConsumerProducer):
print(f"Output Tokens: {outputtokens}", flush=True)
print("Send response...", flush=True)
r = TextCompletionResponse(response=resp, error=None, in_token=inputtokens, out_token=outputtokens, model=self.model)
r = TextCompletionResponse(
response=resp,
error=None,
in_token=inputtokens,
out_token=outputtokens,
model=self.model
)
self.send(r, properties={"id": id})
print("Done.", flush=True)
# FIXME: Wrong exception, don't know what this LLM throws
# for a rate limit
except ResourceExhausted as e:
print("Send rate limit response...", flush=True)
print("Hit rate limit:", e, flush=True)
r = TextCompletionResponse(
error=Error(
type = "rate-limit",
message = str(e),
),
response=None,
in_token=None,
out_token=None,
model=None,
)
self.producer.send(r, properties={"id": id})
self.consumer.acknowledge(msg)
# Leave rate limit retries to the default handler
raise TooManyRequests()
except Exception as e:
print(f"Exception: {e}")
# Apart from rate limits, treat all exceptions as unrecoverable
print(type(e), flush=True)
print(f"Exception: {e}", flush=True)
print("Send error response...", flush=True)

View file

@ -126,26 +126,7 @@ class Processor(ConsumerProducer):
print("Done.", flush=True)
# FIXME: Wrong exception, don't know what this LLM throws
# for a rate limit
except TooManyRequests:
print("Send rate limit response...", flush=True)
r = TextCompletionResponse(
error=Error(
type = "rate-limit",
message = str(e),
),
response=None,
in_token=None,
out_token=None,
model=None,
)
self.producer.send(r, properties={"id": id})
self.consumer.acknowledge(msg)
# SLM, presumably there aren't rate limits
except Exception as e:

View file

@ -100,26 +100,7 @@ class Processor(ConsumerProducer):
print("Done.", flush=True)
# FIXME: Wrong exception, don't know what this LLM throws
# for a rate limit
except TooManyRequests:
print("Send rate limit response...", flush=True)
r = TextCompletionResponse(
error=Error(
type = "rate-limit",
message = str(e),
),
response=None,
in_token=None,
out_token=None,
model=None,
)
self.producer.send(r, properties={"id": id})
self.consumer.acknowledge(msg)
# SLM, presumably no rate limits
except Exception as e:

View file

@ -4,7 +4,7 @@ Simple LLM service, performs text prompt completion using OpenAI.
Input is prompt, output is response.
"""
from openai import OpenAI
from openai import OpenAI, RateLimitError
from prometheus_client import Histogram
import os
@ -87,8 +87,6 @@ class Processor(ConsumerProducer):
try:
# FIXME: Rate limits
with __class__.text_completion_metric.time():
resp = self.openai.chat.completions.create(
@ -134,27 +132,15 @@ class Processor(ConsumerProducer):
# FIXME: Wrong exception, don't know what this LLM throws
# for a rate limit
except TooManyRequests:
except openai.RateLimitError:
print("Send rate limit response...", flush=True)
r = TextCompletionResponse(
error=Error(
type = "rate-limit",
message = str(e),
),
response=None,
in_token=None,
out_token=None,
model=None,
)
self.producer.send(r, properties={"id": id})
self.consumer.acknowledge(msg)
# Leave rate limit retries to the base handler
raise TooManyRequests()
except Exception as e:
# Apart from rate limits, treat all exceptions as unrecoverable
print(f"Exception: {e}")
print("Send error response...", flush=True)

View file

@ -30,6 +30,8 @@ class Processor(ConsumerProducer):
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 | {
@ -39,10 +41,11 @@ class Processor(ConsumerProducer):
"input_schema": DocumentEmbeddingsRequest,
"output_schema": DocumentEmbeddingsResponse,
"store_uri": store_uri,
"api_key": api_key,
}
)
self.client = QdrantClient(url=store_uri)
self.client = QdrantClient(url=store_uri, api_key=api_key)
def handle(self, msg):
@ -111,7 +114,13 @@ class Processor(ConsumerProducer):
parser.add_argument(
'-t', '--store-uri',
default=default_store_uri,
help=f'Milvus store URI (default: {default_store_uri})'
help=f'Qdrant store URI (default: {default_store_uri})'
)
parser.add_argument(
'-k', '--api-key',
default=None,
help=f'API key for qdrant (default: None)'
)
def run():

View file

@ -30,6 +30,7 @@ class Processor(ConsumerProducer):
output_queue = params.get("output_queue", default_output_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 | {
@ -39,10 +40,11 @@ class Processor(ConsumerProducer):
"input_schema": GraphEmbeddingsRequest,
"output_schema": GraphEmbeddingsResponse,
"store_uri": store_uri,
"api_key": api_key,
}
)
self.client = QdrantClient(url=store_uri)
self.client = QdrantClient(url=store_uri, api_key=api_key)
def create_value(self, ent):
if ent.startswith("http://") or ent.startswith("https://"):
@ -137,7 +139,13 @@ class Processor(ConsumerProducer):
parser.add_argument(
'-t', '--store-uri',
default=default_store_uri,
help=f'Milvus store URI (default: {default_store_uri})'
help=f'Qdrant store URI (default: {default_store_uri})'
)
parser.add_argument(
'-k', '--api-key',
default=None,
help=f'API key for qdrant (default: None)'
)
def run():

View file

@ -26,6 +26,8 @@ class Processor(ConsumerProducer):
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 | {
@ -35,10 +37,14 @@ class Processor(ConsumerProducer):
"input_schema": TriplesQueryRequest,
"output_schema": TriplesQueryResponse,
"graph_host": graph_host,
"graph_username": graph_username,
"graph_password": graph_password,
}
)
self.graph_host = [graph_host]
self.username = graph_username
self.password = graph_password
self.table = None
def create_value(self, ent):
@ -56,10 +62,17 @@ class Processor(ConsumerProducer):
table = (v.user, v.collection)
if table != self.table:
self.tg = TrustGraph(
hosts=self.graph_host,
keyspace=v.user, table=v.collection,
)
if self.username and self.password:
self.tg = TrustGraph(
hosts=self.graph_host,
keyspace=v.user, table=v.collection,
username=self.username, password=self.password
)
else:
self.tg = TrustGraph(
hosts=self.graph_host,
keyspace=v.user, table=v.collection,
)
self.table = table
# Sender-produced ID
@ -176,6 +189,19 @@ class Processor(ConsumerProducer):
default="localhost",
help=f'Graph host (default: localhost)'
)
parser.add_argument(
'--graph-username',
default=None,
help=f'Cassandra username'
)
parser.add_argument(
'--graph-password',
default=None,
help=f'Cassandra password'
)
def run():

View file

@ -3,15 +3,16 @@
Accepts entity/vector pairs and writes them to a Milvus store.
"""
from .... schema import ChunkEmbeddings
from .... schema import chunk_embeddings_ingest_queue
from .... log_level import LogLevel
from .... direct.milvus_doc_embeddings import DocVectors
from .... schema import DocumentEmbeddings
from .... schema import document_embeddings_store_queue
from .... log_level import LogLevel
from .... base import Consumer
module = ".".join(__name__.split(".")[1:-1])
default_input_queue = chunk_embeddings_ingest_queue
default_input_queue = document_embeddings_store_queue
default_subscriber = module
default_store_uri = 'http://localhost:19530'
@ -27,7 +28,7 @@ class Processor(Consumer):
**params | {
"input_queue": input_queue,
"subscriber": subscriber,
"input_schema": ChunkEmbeddings,
"input_schema": DocumentEmbeddings,
"store_uri": store_uri,
}
)
@ -38,11 +39,16 @@ class Processor(Consumer):
v = msg.value()
chunk = v.chunk.decode("utf-8")
for emb in v.chunks:
if v.chunk != "" and v.chunk is not None:
for vec in v.vectors:
self.vecstore.insert(vec, chunk)
chunk = emb.chunk.decode("utf-8")
if chunk == "" or chunk is None: continue
for vec in emb.vectors:
if chunk != "" and v.chunk is not None:
for vec in v.vectors:
self.vecstore.insert(vec, chunk)
@staticmethod
def add_args(parser):

View file

@ -11,14 +11,14 @@ import time
import uuid
import os
from .... schema import ChunkEmbeddings
from .... schema import chunk_embeddings_ingest_queue
from .... schema import DocumentEmbeddings
from .... schema import document_embeddings_store_queue
from .... log_level import LogLevel
from .... base import Consumer
module = ".".join(__name__.split(".")[1:-1])
default_input_queue = chunk_embeddings_ingest_queue
default_input_queue = document_embeddings_store_queue
default_subscriber = module
default_api_key = os.getenv("PINECONE_API_KEY", "not-specified")
default_cloud = "aws"
@ -54,7 +54,7 @@ class Processor(Consumer):
**params | {
"input_queue": input_queue,
"subscriber": subscriber,
"input_schema": ChunkEmbeddings,
"input_schema": DocumentEmbeddings,
"url": self.url,
}
)
@ -65,71 +65,74 @@ class Processor(Consumer):
v = msg.value()
chunk = v.chunk.decode("utf-8")
for emb in v.chunks:
if chunk == "": return
chunk = emb.chunk.decode("utf-8")
if chunk == "" or chunk is None: continue
for vec in v.vectors:
for vec in emb.vectors:
dim = len(vec)
collection = (
"d-" + v.metadata.user + "-" + str(dim)
)
for vec in v.vectors:
if index_name != self.last_index_name:
dim = len(vec)
collection = (
"d-" + v.metadata.user + "-" + str(dim)
)
if not self.pinecone.has_index(index_name):
if index_name != self.last_index_name:
try:
if not self.pinecone.has_index(index_name):
self.pinecone.create_index(
name = index_name,
dimension = dim,
metric = "cosine",
spec = ServerlessSpec(
cloud = self.cloud,
region = self.region,
)
)
try:
for i in range(0, 1000):
self.pinecone.create_index(
name = index_name,
dimension = dim,
metric = "cosine",
spec = ServerlessSpec(
cloud = self.cloud,
region = self.region,
)
)
if self.pinecone.describe_index(
index_name
).status["ready"]:
break
for i in range(0, 1000):
time.sleep(1)
if self.pinecone.describe_index(
index_name
).status["ready"]:
break
if not self.pinecone.describe_index(
index_name
).status["ready"]:
raise RuntimeError(
"Gave up waiting for index creation"
)
time.sleep(1)
except Exception as e:
print("Pinecone index creation failed")
raise e
if not self.pinecone.describe_index(
index_name
).status["ready"]:
raise RuntimeError(
"Gave up waiting for index creation"
)
print(f"Index {index_name} created", flush=True)
except Exception as e:
print("Pinecone index creation failed")
raise e
self.last_index_name = index_name
print(f"Index {index_name} created", flush=True)
index = self.pinecone.Index(index_name)
self.last_index_name = index_name
records = [
{
"id": id,
"values": vec,
"metadata": { "doc": chunk },
}
]
index = self.pinecone.Index(index_name)
index.upsert(
vectors = records,
namespace = v.metadata.collection,
)
records = [
{
"id": id,
"values": vec,
"metadata": { "doc": chunk },
}
]
index.upsert(
vectors = records,
namespace = v.metadata.collection,
)
@staticmethod
def add_args(parser):

View file

@ -8,14 +8,14 @@ from qdrant_client.models import PointStruct
from qdrant_client.models import Distance, VectorParams
import uuid
from .... schema import ChunkEmbeddings
from .... schema import chunk_embeddings_ingest_queue
from .... schema import DocumentEmbeddings
from .... schema import document_embeddings_store_queue
from .... log_level import LogLevel
from .... base import Consumer
module = ".".join(__name__.split(".")[1:-1])
default_input_queue = chunk_embeddings_ingest_queue
default_input_queue = document_embeddings_store_queue
default_subscriber = module
default_store_uri = 'http://localhost:6333'
@ -26,13 +26,15 @@ class Processor(Consumer):
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": ChunkEmbeddings,
"input_schema": DocumentEmbeddings,
"store_uri": store_uri,
"api_key": api_key,
}
)
@ -44,47 +46,48 @@ class Processor(Consumer):
v = msg.value()
chunk = v.chunk.decode("utf-8")
for emb in v.chunks:
if chunk == "": return
chunk = emb.chunk.decode("utf-8")
if chunk == "": return
for vec in v.vectors:
for vec in emb.vectors:
dim = len(vec)
collection = (
"d_" + v.metadata.user + "_" + v.metadata.collection + "_" +
str(dim)
)
dim = len(vec)
collection = (
"d_" + v.metadata.user + "_" + v.metadata.collection + "_" +
str(dim)
)
if collection != self.last_collection:
if collection != self.last_collection:
if not self.client.collection_exists(collection):
if not self.client.collection_exists(collection):
try:
self.client.create_collection(
collection_name=collection,
vectors_config=VectorParams(
size=dim, distance=Distance.DOT
),
try:
self.client.create_collection(
collection_name=collection,
vectors_config=VectorParams(
size=dim, distance=Distance.COSINE
),
)
except Exception as e:
print("Qdrant collection creation failed")
raise e
self.last_collection = collection
self.client.upsert(
collection_name=collection,
points=[
PointStruct(
id=str(uuid.uuid4()),
vector=vec,
payload={
"doc": chunk,
}
)
except Exception as e:
print("Qdrant collection creation failed")
raise e
self.last_collection = collection
self.client.upsert(
collection_name=collection,
points=[
PointStruct(
id=str(uuid.uuid4()),
vector=vec,
payload={
"doc": chunk,
}
)
]
)
]
)
@staticmethod
def add_args(parser):
@ -96,7 +99,13 @@ class Processor(Consumer):
parser.add_argument(
'-t', '--store-uri',
default=default_store_uri,
help=f'Qdrant store URI (default: {default_store_uri})'
help=f'Qdrant URI (default: {default_store_uri})'
)
parser.add_argument(
'-k', '--api-key',
default=None,
help=f'Qdrant API key (default: None)'
)
def run():

View file

@ -38,9 +38,11 @@ class Processor(Consumer):
v = msg.value()
if v.entity.value != "":
for vec in v.vectors:
self.vecstore.insert(vec, v.entity.value)
for entity in v.entities:
if entity.entity.value != "" and entity.entity.value is not None:
for vec in entity.vectors:
self.vecstore.insert(vec, entity.entity.value)
@staticmethod
def add_args(parser):

View file

@ -60,76 +60,83 @@ class Processor(Consumer):
self.last_index_name = None
def create_index(self, index_name, dim):
self.pinecone.create_index(
name = index_name,
dimension = dim,
metric = "cosine",
spec = ServerlessSpec(
cloud = self.cloud,
region = self.region,
)
)
for i in range(0, 1000):
if self.pinecone.describe_index(
index_name
).status["ready"]:
break
time.sleep(1)
if not self.pinecone.describe_index(
index_name
).status["ready"]:
raise RuntimeError(
"Gave up waiting for index creation"
)
def handle(self, msg):
v = msg.value()
id = str(uuid.uuid4())
if v.entity.value == "" or v.entity.value is None: return
for entity in v.entities:
for vec in v.vectors:
if entity.entity.value == "" or entity.entity.value is None:
continue
dim = len(vec)
for vec in entity.vectors:
index_name = (
"t-" + v.metadata.user + "-" + str(dim)
)
dim = len(vec)
if index_name != self.last_index_name:
index_name = (
"t-" + v.metadata.user + "-" + str(dim)
)
if not self.pinecone.has_index(index_name):
if index_name != self.last_index_name:
try:
if not self.pinecone.has_index(index_name):
self.pinecone.create_index(
name = index_name,
dimension = dim,
metric = "cosine",
spec = ServerlessSpec(
cloud = self.cloud,
region = self.region,
)
)
try:
for i in range(0, 1000):
self.create_index(index_name, dim)
if self.pinecone.describe_index(
index_name
).status["ready"]:
break
except Exception as e:
print("Pinecone index creation failed")
raise e
time.sleep(1)
print(f"Index {index_name} created", flush=True)
if not self.pinecone.describe_index(
index_name
).status["ready"]:
raise RuntimeError(
"Gave up waiting for index creation"
)
self.last_index_name = index_name
except Exception as e:
print("Pinecone index creation failed")
raise e
index = self.pinecone.Index(index_name)
print(f"Index {index_name} created", flush=True)
records = [
{
"id": id,
"values": vec,
"metadata": { "entity": entity.entity.value },
}
]
self.last_index_name = index_name
index = self.pinecone.Index(index_name)
records = [
{
"id": id,
"values": vec,
"metadata": { "entity": v.entity.value },
}
]
index.upsert(
vectors = records,
namespace = v.metadata.collection,
)
index.upsert(
vectors = records,
namespace = v.metadata.collection,
)
@staticmethod
def add_args(parser):

View file

@ -26,6 +26,7 @@ class Processor(Consumer):
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 | {
@ -33,56 +34,67 @@ class Processor(Consumer):
"subscriber": subscriber,
"input_schema": GraphEmbeddings,
"store_uri": store_uri,
"api_key": api_key,
}
)
self.last_collection = None
self.client = QdrantClient(url=store_uri)
self.client = QdrantClient(url=store_uri, api_key=api_key)
def get_collection(self, dim, user, collection):
cname = (
"t_" + user + "_" + collection + "_" + str(dim)
)
if cname != self.last_collection:
if not self.client.collection_exists(cname):
try:
self.client.create_collection(
collection_name=cname,
vectors_config=VectorParams(
size=dim, distance=Distance.COSINE
),
)
except Exception as e:
print("Qdrant collection creation failed")
raise e
self.last_collection = cname
return cname
def handle(self, msg):
v = msg.value()
if v.entity.value == "" or v.entity.value is None: return
for entity in v.entities:
for vec in v.vectors:
if entity.entity.value == "" or entity.entity.value is None: return
dim = len(vec)
collection = (
"t_" + v.metadata.user + "_" + v.metadata.collection + "_" +
str(dim)
)
for vec in entity.vectors:
if collection != self.last_collection:
dim = len(vec)
if not self.client.collection_exists(collection):
collection = self.get_collection(
dim, v.metadata.user, v.metadata.collection
)
try:
self.client.create_collection(
collection_name=collection,
vectors_config=VectorParams(
size=dim, distance=Distance.COSINE
),
self.client.upsert(
collection_name=collection,
points=[
PointStruct(
id=str(uuid.uuid4()),
vector=vec,
payload={
"entity": entity.entity.value,
}
)
except Exception as e:
print("Qdrant collection creation failed")
raise e
self.last_collection = collection
self.client.upsert(
collection_name=collection,
points=[
PointStruct(
id=str(uuid.uuid4()),
vector=vec,
payload={
"entity": v.entity.value,
}
)
]
)
]
)
@staticmethod
def add_args(parser):
@ -96,6 +108,12 @@ class Processor(Consumer):
default=default_store_uri,
help=f'Qdrant store URI (default: {default_store_uri})'
)
parser.add_argument(
'-k', '--api-key',
default=None,
help=f'Qdrant API key'
)
def run():

View file

@ -29,6 +29,8 @@ class Processor(Consumer):
input_queue = params.get("input_queue", default_input_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 | {
@ -36,10 +38,16 @@ class Processor(Consumer):
"subscriber": subscriber,
"input_schema": Rows,
"graph_host": graph_host,
"graph_username": graph_username,
"graph_password": graph_password,
}
)
self.cluster = Cluster(graph_host.split(","))
if graph_username and graph_password:
auth_provider = PlainTextAuthProvider(username=graph_username, password=graph_password)
self.cluster = Cluster(graph_host.split(","), auth_provider=auth_provider)
else:
self.cluster = Cluster(graph_host.split(","))
self.session = self.cluster.connect()
self.tables = set()
@ -120,6 +128,18 @@ class Processor(Consumer):
default="localhost",
help=f'Graph host (default: localhost)'
)
parser.add_argument(
'--graph-username',
default=None,
help=f'Cassandra username'
)
parser.add_argument(
'--graph-password',
default=None,
help=f'Cassandra password'
)
def run():

View file

@ -28,6 +28,8 @@ class Processor(Consumer):
input_queue = params.get("input_queue", default_input_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 | {
@ -35,10 +37,14 @@ class Processor(Consumer):
"subscriber": subscriber,
"input_schema": Triples,
"graph_host": graph_host,
"graph_username": graph_username,
"graph_password": graph_password,
}
)
self.graph_host = [graph_host]
self.username = graph_username
self.password = graph_password
self.table = None
def handle(self, msg):
@ -52,10 +58,17 @@ class Processor(Consumer):
self.tg = None
try:
self.tg = TrustGraph(
hosts=self.graph_host,
keyspace=v.metadata.user, table=v.metadata.collection,
)
if self.username and self.password:
self.tg = TrustGraph(
hosts=self.graph_host,
keyspace=v.metadata.user, table=v.metadata.collection,
username=self.username, password=self.password
)
else:
self.tg = TrustGraph(
hosts=self.graph_host,
keyspace=v.metadata.user, table=v.metadata.collection,
)
except Exception as e:
print("Exception", e, flush=True)
time.sleep(1)
@ -82,6 +95,18 @@ class Processor(Consumer):
default="localhost",
help=f'Graph host (default: localhost)'
)
parser.add_argument(
'--graph-username',
default=None,
help=f'Cassandra username'
)
parser.add_argument(
'--graph-password',
default=None,
help=f'Cassandra password'
)
def run():

View file

@ -55,6 +55,14 @@ class Processor(Consumer):
def create_indexes(self, session):
# Race condition, index creation failure is ignored. Right thing
# to do if the index already exists. Wrong thing to do if it's
# because the store is not up yet
# In real-world cases, Memgraph will start up quicker than Pulsar
# and this process will restart several times until Pulsar arrives,
# so should be safe
print("Create indexes...", flush=True)
try:

View file

@ -50,6 +50,50 @@ class Processor(Consumer):
self.io = GraphDatabase.driver(graph_host, auth=(username, password))
with self.io.session(database=self.db) as session:
self.create_indexes(session)
def create_indexes(self, session):
# Race condition, index creation failure is ignored. Right thing
# to do if the index already exists. Wrong thing to do if it's
# because the store is not up yet
# In real-world cases, Neo4j will start up quicker than Pulsar
# and this process will restart several times until Pulsar arrives,
# so should be safe
print("Create indexes...", flush=True)
try:
session.run(
"CREATE INDEX Node_uri FOR (n:Node) ON (n.uri)",
)
except Exception as e:
print(e, flush=True)
# Maybe index already exists
print("Index create failure ignored", flush=True)
try:
session.run(
"CREATE INDEX Literal_value FOR (n:Literal) ON (n.value)",
)
except Exception as e:
print(e, flush=True)
# Maybe index already exists
print("Index create failure ignored", flush=True)
try:
session.run(
"CREATE INDEX Rel_uri FOR ()-[r:Rel]-() ON (r.uri)",
)
except Exception as e:
print(e, flush=True)
# Maybe index already exists
print("Index create failure ignored", flush=True)
print("Index creation done", flush=True)
def create_node(self, uri):
print("Create node", uri)