Add pulsar API token check

This commit is contained in:
Tyler O 2025-02-10 16:30:53 +00:00
parent d0ae772fd6
commit a5d5b4ca4a
56 changed files with 319 additions and 82 deletions

View file

@ -11,6 +11,7 @@ from .. log_level import LogLevel
class BaseProcessor: class BaseProcessor:
default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://pulsar:6650') default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://pulsar:6650')
default_pulsar_api_key = os.getenv("PULSAR_API_KEY", None)
def __init__(self, **params): def __init__(self, **params):
@ -28,14 +29,22 @@ class BaseProcessor:
}) })
pulsar_host = params.get("pulsar_host", self.default_pulsar_host) pulsar_host = params.get("pulsar_host", self.default_pulsar_host)
pulsar_api_key = params.get("pulsar_api_key", None)
log_level = params.get("log_level", LogLevel.INFO) log_level = params.get("log_level", LogLevel.INFO)
self.pulsar_host = pulsar_host self.pulsar_host = pulsar_host
if pulsar_api_key:
self.client = pulsar.Client( auth = pulsar.AuthenticationToken(pulsar_api_key)
self.client = pulsar.Client(
pulsar_host,
authentication=auth,
logger=pulsar.ConsoleLogger(log_level.to_pulsar())
)
else:
self.client = pulsar.Client(
pulsar_host, pulsar_host,
logger=pulsar.ConsoleLogger(log_level.to_pulsar()) logger=pulsar.ConsoleLogger(log_level.to_pulsar())
) )
def __del__(self): def __del__(self):
@ -51,6 +60,12 @@ class BaseProcessor:
default=__class__.default_pulsar_host, default=__class__.default_pulsar_host,
help=f'Pulsar host (default: {__class__.default_pulsar_host})', help=f'Pulsar host (default: {__class__.default_pulsar_host})',
) )
parser.add_argument(
'--pulsar-api-key',
default=__class__.default_pulsar_api_key,
help=f'Pulsar API key',
)
parser.add_argument( parser.add_argument(
'-l', '--log-level', '-l', '--log-level',

View file

@ -20,6 +20,7 @@ class AgentClient(BaseClient):
input_queue=None, input_queue=None,
output_queue=None, output_queue=None,
pulsar_host="pulsar://pulsar:6650", pulsar_host="pulsar://pulsar:6650",
pulsar_api_key=None,
): ):
if input_queue is None: input_queue = agent_request_queue if input_queue is None: input_queue = agent_request_queue
@ -33,6 +34,7 @@ class AgentClient(BaseClient):
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
input_schema=AgentRequest, input_schema=AgentRequest,
output_schema=AgentResponse, output_schema=AgentResponse,
pulsar_api_key=pulsar_api_key
) )
def request( def request(

View file

@ -27,6 +27,7 @@ class BaseClient:
input_schema=None, input_schema=None,
output_schema=None, output_schema=None,
pulsar_host="pulsar://pulsar:6650", pulsar_host="pulsar://pulsar:6650",
pulsar_api_key=None,
): ):
if input_queue == None: raise RuntimeError("Need input_queue") if input_queue == None: raise RuntimeError("Need input_queue")
@ -37,10 +38,18 @@ class BaseClient:
if subscriber == None: if subscriber == None:
subscriber = str(uuid.uuid4()) subscriber = str(uuid.uuid4())
self.client = pulsar.Client( if pulsar_api_key:
auth = pulsar.AuthenticationToken(pulsar_api_key)
self.client = pulsar.Client(
pulsar_host, pulsar_host,
logger=pulsar.ConsoleLogger(log_level), authentication=auth,
) logger=pulsar.ConsoleLogger(log_level.to_pulsar())
)
else:
self.client = pulsar.Client(
pulsar_host,
logger=pulsar.ConsoleLogger(log_level.to_pulsar())
)
self.producer = self.client.create_producer( self.producer = self.client.create_producer(
topic=input_queue, topic=input_queue,

View file

@ -20,6 +20,7 @@ class DocumentEmbeddingsClient(BaseClient):
input_queue=None, input_queue=None,
output_queue=None, output_queue=None,
pulsar_host="pulsar://pulsar:6650", pulsar_host="pulsar://pulsar:6650",
pulsar_api_key=None,
): ):
if input_queue == None: if input_queue == None:
@ -34,6 +35,7 @@ class DocumentEmbeddingsClient(BaseClient):
input_queue=input_queue, input_queue=input_queue,
output_queue=output_queue, output_queue=output_queue,
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
input_schema=DocumentEmbeddingsRequest, input_schema=DocumentEmbeddingsRequest,
output_schema=DocumentEmbeddingsResponse, output_schema=DocumentEmbeddingsResponse,
) )

View file

@ -20,6 +20,7 @@ class DocumentRagClient(BaseClient):
input_queue=None, input_queue=None,
output_queue=None, output_queue=None,
pulsar_host="pulsar://pulsar:6650", pulsar_host="pulsar://pulsar:6650",
pulsar_api_key=None,
): ):
if input_queue == None: if input_queue == None:
@ -34,6 +35,7 @@ class DocumentRagClient(BaseClient):
input_queue=input_queue, input_queue=input_queue,
output_queue=output_queue, output_queue=output_queue,
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
input_schema=DocumentRagQuery, input_schema=DocumentRagQuery,
output_schema=DocumentRagResponse, output_schema=DocumentRagResponse,
) )

View file

@ -20,6 +20,7 @@ class EmbeddingsClient(BaseClient):
output_queue=None, output_queue=None,
subscriber=None, subscriber=None,
pulsar_host="pulsar://pulsar:6650", pulsar_host="pulsar://pulsar:6650",
pulsar_api_key=None,
): ):
if input_queue == None: if input_queue == None:
@ -34,6 +35,7 @@ class EmbeddingsClient(BaseClient):
input_queue=input_queue, input_queue=input_queue,
output_queue=output_queue, output_queue=output_queue,
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
input_schema=EmbeddingsRequest, input_schema=EmbeddingsRequest,
output_schema=EmbeddingsResponse, output_schema=EmbeddingsResponse,
) )

View file

@ -20,6 +20,7 @@ class GraphEmbeddingsClient(BaseClient):
input_queue=None, input_queue=None,
output_queue=None, output_queue=None,
pulsar_host="pulsar://pulsar:6650", pulsar_host="pulsar://pulsar:6650",
pulsar_api_key=None,
): ):
if input_queue == None: if input_queue == None:
@ -34,6 +35,7 @@ class GraphEmbeddingsClient(BaseClient):
input_queue=input_queue, input_queue=input_queue,
output_queue=output_queue, output_queue=output_queue,
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
input_schema=GraphEmbeddingsRequest, input_schema=GraphEmbeddingsRequest,
output_schema=GraphEmbeddingsResponse, output_schema=GraphEmbeddingsResponse,
) )

View file

@ -20,6 +20,7 @@ class GraphRagClient(BaseClient):
input_queue=None, input_queue=None,
output_queue=None, output_queue=None,
pulsar_host="pulsar://pulsar:6650", pulsar_host="pulsar://pulsar:6650",
pulsar_api_key=None,
): ):
if input_queue == None: if input_queue == None:
@ -34,6 +35,7 @@ class GraphRagClient(BaseClient):
input_queue=input_queue, input_queue=input_queue,
output_queue=output_queue, output_queue=output_queue,
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
input_schema=GraphRagQuery, input_schema=GraphRagQuery,
output_schema=GraphRagResponse, output_schema=GraphRagResponse,
) )

View file

@ -20,6 +20,7 @@ class LlmClient(BaseClient):
input_queue=None, input_queue=None,
output_queue=None, output_queue=None,
pulsar_host="pulsar://pulsar:6650", pulsar_host="pulsar://pulsar:6650",
pulsar_api_key=None,
): ):
if input_queue is None: input_queue = text_completion_request_queue if input_queue is None: input_queue = text_completion_request_queue
@ -31,6 +32,7 @@ class LlmClient(BaseClient):
input_queue=input_queue, input_queue=input_queue,
output_queue=output_queue, output_queue=output_queue,
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
input_schema=TextCompletionRequest, input_schema=TextCompletionRequest,
output_schema=TextCompletionResponse, output_schema=TextCompletionResponse,
) )

View file

@ -39,6 +39,7 @@ class PromptClient(BaseClient):
input_queue=None, input_queue=None,
output_queue=None, output_queue=None,
pulsar_host="pulsar://pulsar:6650", pulsar_host="pulsar://pulsar:6650",
pulsar_api_key=None,
): ):
if input_queue == None: if input_queue == None:
@ -53,6 +54,7 @@ class PromptClient(BaseClient):
input_queue=input_queue, input_queue=input_queue,
output_queue=output_queue, output_queue=output_queue,
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
input_schema=PromptRequest, input_schema=PromptRequest,
output_schema=PromptResponse, output_schema=PromptResponse,
) )

View file

@ -21,6 +21,7 @@ class TriplesQueryClient(BaseClient):
input_queue=None, input_queue=None,
output_queue=None, output_queue=None,
pulsar_host="pulsar://pulsar:6650", pulsar_host="pulsar://pulsar:6650",
pulsar_api_key=None,
): ):
if input_queue == None: if input_queue == None:
@ -34,6 +35,7 @@ class TriplesQueryClient(BaseClient):
subscriber=subscriber, subscriber=subscriber,
input_queue=input_queue, input_queue=input_queue,
output_queue=output_queue, output_queue=output_queue,
pulsar_api_key=pulsar_api_key,
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
input_schema=TriplesQueryRequest, input_schema=TriplesQueryRequest,
output_schema=TriplesQueryResponse, output_schema=TriplesQueryResponse,

View file

@ -9,12 +9,13 @@ import os
from trustgraph.clients.triples_query_client import TriplesQueryClient from trustgraph.clients.triples_query_client import TriplesQueryClient
default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650') default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650')
default_pulsar_api_key = os.getenv("PULSAR_API_KEY", None)
default_user = 'trustgraph' default_user = 'trustgraph'
default_collection = 'default' default_collection = 'default'
def show_graph(pulsar, user, collection): def show_graph(pulsar, user, collection, pulsar_api_key=None):
tq = TriplesQueryClient(pulsar_host=pulsar) tq = TriplesQueryClient(pulsar_host=pulsar, pulsar_api_key=pulsar_api_key)
rows = tq.request( rows = tq.request(
user=user, collection=collection, user=user, collection=collection,
@ -48,7 +49,13 @@ def main():
default=default_collection, default=default_collection,
help=f'Collection ID (default: {default_collection})' help=f'Collection ID (default: {default_collection})'
) )
parser.add_argument(
'--pulsar-api-key',
default=default_pulsar_api_key,
help=f'Pulsar API key',
)
args = parser.parse_args() args = parser.parse_args()
try: try:
@ -56,6 +63,7 @@ def main():
show_graph( show_graph(
pulsar=args.pulsar_host, user=args.user, pulsar=args.pulsar_host, user=args.user,
collection=args.collection, collection=args.collection,
pulsar_api_key=args.pulsar_api_key,
) )
except Exception as e: except Exception as e:

View file

@ -13,10 +13,11 @@ import io
import sys import sys
default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650') default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650')
default_pulsar_api_key = os.getenv("PULSAR_API_KEY", None)
def show_graph(pulsar): def show_graph(pulsar, pulsar_api_key=None):
tq = TriplesQueryClient(pulsar_host=pulsar) tq = TriplesQueryClient(pulsar_host=pulsar, pulsar_api_key=pulsar_api_key)
rows = tq.request(None, None, None, limit=10_000_000) rows = tq.request(None, None, None, limit=10_000_000)
@ -60,12 +61,18 @@ def main():
default=default_pulsar_host, default=default_pulsar_host,
help=f'Pulsar host (default: {default_pulsar_host})', help=f'Pulsar host (default: {default_pulsar_host})',
) )
parser.add_argument(
'--pulsar-api-key',
default=default_pulsar_api_key,
help=f'Pulsar API key',
)
args = parser.parse_args() args = parser.parse_args()
try: try:
show_graph(args.pulsar_host) show_graph(args.pulsar_host, pulsar_api_key=args.pulsar_api_key)
except Exception as e: except Exception as e:

View file

@ -11,6 +11,7 @@ import textwrap
from trustgraph.clients.agent_client import AgentClient from trustgraph.clients.agent_client import AgentClient
default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650') default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650')
default_pulsar_api_key = os.getenv("PULSAR_API_KEY", None)
default_user = 'trustgraph' default_user = 'trustgraph'
default_collection = 'default' default_collection = 'default'
@ -29,10 +30,10 @@ def output(text, prefix="> ", width=78):
def query( def query(
pulsar_host, query, user, collection, pulsar_host, query, user, collection,
plan=None, state=None, verbose=False plan=None, state=None, verbose=False, pulsar_api_key=None
): ):
am = AgentClient(pulsar_host=pulsar_host) am = AgentClient(pulsar_host=pulsar_host, pulsar_api_key=pulsar_api_key)
if verbose: if verbose:
output(wrap(query), "\U00002753 ") output(wrap(query), "\U00002753 ")
@ -100,6 +101,12 @@ def main():
action="store_true", action="store_true",
help=f'Output thinking/observations' help=f'Output thinking/observations'
) )
parser.add_argument(
'--pulsar-api-key',
default=default_pulsar_api_key,
help=f'Pulsar API key',
)
args = parser.parse_args() args = parser.parse_args()
@ -113,6 +120,7 @@ def main():
plan=args.plan, plan=args.plan,
state=args.state, state=args.state,
verbose=args.verbose, verbose=args.verbose,
pulsar_api_key=args.pulsar_api_key,
) )
except Exception as e: except Exception as e:

View file

@ -11,10 +11,12 @@ import json
from trustgraph.clients.llm_client import LlmClient from trustgraph.clients.llm_client import LlmClient
default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650') default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650')
default_pulsar_api_key = os.getenv("PULSAR_API_KEY", None)
def query(pulsar_host, system, prompt):
cli = LlmClient(pulsar_host=pulsar_host) def query(pulsar_host, system, prompt, pulsar_api_key=None):
cli = LlmClient(pulsar_host=pulsar_host, pulsar_api_key=pulsar_api_key)
resp = cli.request(system=system, prompt=prompt) resp = cli.request(system=system, prompt=prompt)
@ -32,7 +34,7 @@ def main():
default=default_pulsar_host, default=default_pulsar_host,
help=f'Pulsar host (default: {default_pulsar_host})', help=f'Pulsar host (default: {default_pulsar_host})',
) )
parser.add_argument( parser.add_argument(
'system', 'system',
nargs=1, nargs=1,
@ -44,6 +46,13 @@ def main():
nargs=1, nargs=1,
help='LLM prompt e.g. What is 2 + 2?', help='LLM prompt e.g. What is 2 + 2?',
) )
parser.add_argument(
'--pulsar-api-key',
default=default_pulsar_api_key,
help=f'Pulsar API key',
)
args = parser.parse_args() args = parser.parse_args()
@ -53,6 +62,7 @@ def main():
pulsar_host=args.pulsar_host, pulsar_host=args.pulsar_host,
system=args.system[0], system=args.system[0],
prompt=args.prompt[0], prompt=args.prompt[0],
pulsar_api_key=args.pulsar_api_key,
) )
except Exception as e: except Exception as e:

View file

@ -15,10 +15,12 @@ import json
from trustgraph.clients.prompt_client import PromptClient from trustgraph.clients.prompt_client import PromptClient
default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650') default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650')
default_pulsar_api_key = os.getenv("PULSAR_API_KEY", None)
def query(pulsar_host, template_id, variables):
cli = PromptClient(pulsar_host=pulsar_host) def query(pulsar_host, template_id, variables, pulsar_api_key=None):
cli = PromptClient(pulsar_host=pulsar_host, pulsar_api_key=pulsar_api_key)
resp = cli.request(id=template_id, variables=variables) resp = cli.request(id=template_id, variables=variables)
@ -55,6 +57,13 @@ def main():
specified multiple times''', specified multiple times''',
) )
parser.add_argument(
'--pulsar-api-key',
default=default_pulsar_api_key,
help=f'Pulsar API key',
)
args = parser.parse_args() args = parser.parse_args()
variables = {} variables = {}
@ -73,6 +82,7 @@ specified multiple times''',
pulsar_host=args.pulsar_host, pulsar_host=args.pulsar_host,
template_id=args.id[0], template_id=args.id[0],
variables=variables, variables=variables,
pulsar_api_key=args.pulsar_api_key,
) )
except Exception as e: except Exception as e:

View file

@ -34,13 +34,22 @@ class Loader:
collection, collection,
log_level, log_level,
metadata, metadata,
pulsar_api_key=None,
): ):
self.client = pulsar.Client( if pulsar_api_key:
auth = pulsar.AuthenticationToken(pulsar_api_key)
self.client = pulsar.Client(
pulsar_host,
authentication=auth,
logger=pulsar.ConsoleLogger(log_level.to_pulsar())
)
else:
self.client = pulsar.Client(
pulsar_host, pulsar_host,
logger=pulsar.ConsoleLogger(log_level.to_pulsar()) logger=pulsar.ConsoleLogger(log_level.to_pulsar())
) )
self.producer = self.client.create_producer( self.producer = self.client.create_producer(
topic=output_queue, topic=output_queue,
schema=JsonSchema(Document), schema=JsonSchema(Document),
@ -120,6 +129,7 @@ def main():
) )
default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650') default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650')
default_pulsar_api_key = os.getenv("PULSAR_API_KEY", None)
default_output_queue = document_ingest_queue default_output_queue = document_ingest_queue
parser.add_argument( parser.add_argument(
@ -127,6 +137,12 @@ def main():
default=default_pulsar_host, default=default_pulsar_host,
help=f'Pulsar host (default: {default_pulsar_host})', help=f'Pulsar host (default: {default_pulsar_host})',
) )
parser.add_argument(
'--pulsar-api-key',
default=default_pulsar_api_key,
help=f'Pulsar API key',
)
parser.add_argument( parser.add_argument(
'-o', '--output-queue', '-o', '--output-queue',
@ -240,6 +256,7 @@ def main():
p = Loader( p = Loader(
pulsar_host=args.pulsar_host, pulsar_host=args.pulsar_host,
pulsar_api_key=args.pulsar_api_key,
output_queue=args.output_queue, output_queue=args.output_queue,
user=args.user, user=args.user,
collection=args.collection, collection=args.collection,

View file

@ -33,12 +33,20 @@ class Loader:
collection, collection,
log_level, log_level,
metadata, metadata,
pulsar_api_key=None,
): ):
if pulsar_api_key:
self.client = pulsar.Client( auth = pulsar.AuthenticationToken(pulsar_api_key)
pulsar_host, self.client = pulsar.Client(
logger=pulsar.ConsoleLogger(log_level.to_pulsar()) pulsar_host,
) authentication=auth,
logger=pulsar.ConsoleLogger(log_level.to_pulsar())
)
else:
self.client = pulsar.Client(
pulsar_host,
logger=pulsar.ConsoleLogger(log_level.to_pulsar())
)
self.producer = self.client.create_producer( self.producer = self.client.create_producer(
topic=output_queue, topic=output_queue,
@ -119,6 +127,8 @@ def main():
) )
default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650') default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650')
default_pulsar_api_key = os.getenv("PULSAR_API_KEY", None)
default_output_queue = text_ingest_queue default_output_queue = text_ingest_queue
parser.add_argument( parser.add_argument(
@ -126,6 +136,12 @@ def main():
default=default_pulsar_host, default=default_pulsar_host,
help=f'Pulsar host (default: {default_pulsar_host})', help=f'Pulsar host (default: {default_pulsar_host})',
) )
parser.add_argument(
'--pulsar-api-key',
default=default_pulsar_api_key,
help=f'Pulsar API key',
)
parser.add_argument( parser.add_argument(
'-o', '--output-queue', '-o', '--output-queue',
@ -239,6 +255,7 @@ def main():
p = Loader( p = Loader(
pulsar_host=args.pulsar_host, pulsar_host=args.pulsar_host,
pulsar_api_key=args.pulsar_api_key,
output_queue=args.output_queue, output_queue=args.output_queue,
user=args.user, user=args.user,
collection=args.collection, collection=args.collection,

View file

@ -19,6 +19,8 @@ from trustgraph.log_level import LogLevel
default_user = 'trustgraph' default_user = 'trustgraph'
default_collection = 'default' default_collection = 'default'
default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650') default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650')
default_pulsar_api_key = os.getenv("PULSAR_API_KEY", None)
default_output_queue = triples_store_queue default_output_queue = triples_store_queue
class Loader: class Loader:
@ -31,12 +33,21 @@ class Loader:
files, files,
user, user,
collection, collection,
pulsar_api_key=None,
): ):
self.client = pulsar.Client( if pulsar_api_key:
pulsar_host, auth = pulsar.AuthenticationToken(pulsar_api_key)
logger=pulsar.ConsoleLogger(log_level.to_pulsar()) self.client = pulsar.Client(
) pulsar_host,
authentication=auth,
logger=pulsar.ConsoleLogger(log_level.to_pulsar())
)
else:
self.client = pulsar.Client(
pulsar_host,
logger=pulsar.ConsoleLogger(log_level.to_pulsar())
)
self.producer = self.client.create_producer( self.producer = self.client.create_producer(
topic=output_queue, topic=output_queue,
@ -98,6 +109,12 @@ def main():
default=default_pulsar_host, default=default_pulsar_host,
help=f'Pulsar host (default: {default_pulsar_host})', help=f'Pulsar host (default: {default_pulsar_host})',
) )
parser.add_argument(
'--pulsar-api-key',
default=default_pulsar_api_key,
help=f'Pulsar API key',
)
parser.add_argument( parser.add_argument(
'-o', '--output-queue', '-o', '--output-queue',
@ -137,6 +154,7 @@ def main():
try: try:
p = Loader( p = Loader(
pulsar_host=args.pulsar_host, pulsar_host=args.pulsar_host,
pulsar_api_key=args.pulsar_api_key,
output_queue=args.output_queue, output_queue=args.output_queue,
log_level=args.log_level, log_level=args.log_level,
files=args.files, files=args.files,

View file

@ -9,12 +9,14 @@ import os
from trustgraph.clients.document_rag_client import DocumentRagClient from trustgraph.clients.document_rag_client import DocumentRagClient
default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650') default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650')
default_pulsar_api_key = os.getenv("PULSAR_API_KEY", None)
default_user = 'trustgraph' default_user = 'trustgraph'
default_collection = 'default' default_collection = 'default'
def query(pulsar_host, query, user, collection): def query(pulsar_host, query, user, collection, pulsar_api_key=None):
rag = DocumentRagClient(pulsar_host=pulsar) rag = DocumentRagClient(pulsar_host=pulsar_host, pulsar_api_key=pulsar_api_key)
resp = rag.request(user=user, collection=collection, query=query) resp = rag.request(user=user, collection=collection, query=query)
print(resp) print(resp)
@ -30,7 +32,12 @@ def main():
default=default_pulsar_host, default=default_pulsar_host,
help=f'Pulsar host (default: {default_pulsar_host})', help=f'Pulsar host (default: {default_pulsar_host})',
) )
parser.add_argument(
'--pulsar-api-key',
default=default_pulsar_api_key,
help=f'Pulsar API key',
)
parser.add_argument( parser.add_argument(
'-q', '--query', '-q', '--query',
required=True, required=True,
@ -55,6 +62,7 @@ def main():
query( query(
pulsar_host=args.pulsar_host, pulsar_host=args.pulsar_host,
pulsar_api_key=args.pulsar_api_key,
query=args.query, query=args.query,
user=args.user, user=args.user,
collection=args.collection, collection=args.collection,

View file

@ -9,12 +9,14 @@ import os
from trustgraph.clients.graph_rag_client import GraphRagClient from trustgraph.clients.graph_rag_client import GraphRagClient
default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650') default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://localhost:6650')
default_pulsar_api_key = os.getenv("PULSAR_API_KEY", None)
default_user = 'trustgraph' default_user = 'trustgraph'
default_collection = 'default' default_collection = 'default'
def query(pulsar_host, query, user, collection): def query(pulsar_host, query, user, collection, pulsar_api_key=None):
rag = GraphRagClient(pulsar_host=pulsar_host) rag = GraphRagClient(pulsar_host=pulsar_host, pulsar_api_key=pulsar_api_key)
resp = rag.request(user=user, collection=collection, query=query) resp = rag.request(user=user, collection=collection, query=query)
print(resp) print(resp)
@ -31,6 +33,12 @@ def main():
help=f'Pulsar host (default: {default_pulsar_host})', help=f'Pulsar host (default: {default_pulsar_host})',
) )
parser.add_argument(
'--pulsar-api-key',
default=default_pulsar_api_key,
help=f'Pulsar API key',
)
parser.add_argument( parser.add_argument(
'-q', '--query', '-q', '--query',
required=True, required=True,
@ -55,6 +63,7 @@ def main():
query( query(
pulsar_host=args.pulsar_host, pulsar_host=args.pulsar_host,
pulsar_api_key=args.pulsar_api_key,
query=args.query, query=args.query,
user=args.user, user=args.user,
collection=args.collection, collection=args.collection,

View file

@ -166,21 +166,24 @@ class Processor(ConsumerProducer):
subscriber=subscriber, subscriber=subscriber,
input_queue=prompt_request_queue, input_queue=prompt_request_queue,
output_queue=prompt_response_queue, output_queue=prompt_response_queue,
pulsar_host = self.pulsar_host pulsar_host = self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
) )
self.llm = LlmClient( self.llm = LlmClient(
subscriber=subscriber, subscriber=subscriber,
input_queue=text_completion_request_queue, input_queue=text_completion_request_queue,
output_queue=text_completion_response_queue, output_queue=text_completion_response_queue,
pulsar_host = self.pulsar_host pulsar_host = self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
) )
self.graph_rag = GraphRagClient( self.graph_rag = GraphRagClient(
subscriber=subscriber, subscriber=subscriber,
input_queue=graph_rag_request_queue, input_queue=graph_rag_request_queue,
output_queue=graph_rag_response_queue, output_queue=graph_rag_response_queue,
pulsar_host = self.pulsar_host pulsar_host = self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
) )
# Need to be able to feed requests to myself # Need to be able to feed requests to myself

View file

@ -21,6 +21,7 @@ class DocumentRag:
def __init__( def __init__(
self, self,
pulsar_host="pulsar://pulsar:6650", pulsar_host="pulsar://pulsar:6650",
pulsar_api_key=None,
pr_request_queue=None, pr_request_queue=None,
pr_response_queue=None, pr_response_queue=None,
emb_request_queue=None, emb_request_queue=None,
@ -62,6 +63,7 @@ class DocumentRag:
subscriber=module + "-de", subscriber=module + "-de",
input_queue=de_request_queue, input_queue=de_request_queue,
output_queue=de_response_queue, output_queue=de_response_queue,
pulsar_api_key=pulsar_api_key,
) )
self.embeddings = EmbeddingsClient( self.embeddings = EmbeddingsClient(
@ -69,6 +71,7 @@ class DocumentRag:
input_queue=emb_request_queue, input_queue=emb_request_queue,
output_queue=emb_response_queue, output_queue=emb_response_queue,
subscriber=module + "-emb", subscriber=module + "-emb",
pulsar_api_key=pulsar_api_key,
) )
self.lang = PromptClient( self.lang = PromptClient(
@ -76,6 +79,7 @@ class DocumentRag:
input_queue=pr_request_queue, input_queue=pr_request_queue,
output_queue=pr_response_queue, output_queue=pr_response_queue,
subscriber=module + "-de-prompt", subscriber=module + "-de-prompt",
pulsar_api_key=pulsar_api_key,
) )
if self.verbose: if self.verbose:

View file

@ -45,6 +45,7 @@ class Processor(ConsumerProducer):
self.embeddings = EmbeddingsClient( self.embeddings = EmbeddingsClient(
pulsar_host=self.pulsar_host, pulsar_host=self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
input_queue=emb_request_queue, input_queue=emb_request_queue,
output_queue=emb_response_queue, output_queue=emb_response_queue,
subscriber=module + "-emb", subscriber=module + "-emb",

View file

@ -54,6 +54,7 @@ class Processor(ConsumerProducer):
self.prompt = PromptClient( self.prompt = PromptClient(
pulsar_host=self.pulsar_host, pulsar_host=self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
input_queue=pr_request_queue, input_queue=pr_request_queue,
output_queue=pr_response_queue, output_queue=pr_response_queue,
subscriber = module + "-prompt", subscriber = module + "-prompt",

View file

@ -76,6 +76,7 @@ class Processor(ConsumerProducer):
self.prompt = PromptClient( self.prompt = PromptClient(
pulsar_host=self.pulsar_host, pulsar_host=self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
input_queue=pr_request_queue, input_queue=pr_request_queue,
output_queue=pr_response_queue, output_queue=pr_response_queue,
subscriber = module + "-prompt", subscriber = module + "-prompt",

View file

@ -52,6 +52,7 @@ class Processor(ConsumerProducer):
self.prompt = PromptClient( self.prompt = PromptClient(
pulsar_host=self.pulsar_host, pulsar_host=self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
input_queue=pr_request_queue, input_queue=pr_request_queue,
output_queue=pr_response_queue, output_queue=pr_response_queue,
subscriber = module + "-prompt", subscriber = module + "-prompt",

View file

@ -112,6 +112,7 @@ class Processor(ConsumerProducer):
self.prompt = PromptClient( self.prompt = PromptClient(
pulsar_host=self.pulsar_host, pulsar_host=self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
input_queue=pr_request_queue, input_queue=pr_request_queue,
output_queue=pr_response_queue, output_queue=pr_response_queue,
subscriber = module + "-prompt", subscriber = module + "-prompt",

View file

@ -7,10 +7,11 @@ from . endpoint import ServiceEndpoint
from . requestor import ServiceRequestor from . requestor import ServiceRequestor
class AgentRequestor(ServiceRequestor): class AgentRequestor(ServiceRequestor):
def __init__(self, pulsar_host, timeout, auth): def __init__(self, pulsar_host, timeout, auth, pulsar_api_key=None):
super(AgentRequestor, self).__init__( super(AgentRequestor, self).__init__(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
request_queue=agent_request_queue, request_queue=agent_request_queue,
response_queue=agent_response_queue, response_queue=agent_response_queue,
request_schema=AgentRequest, request_schema=AgentRequest,

View file

@ -7,10 +7,11 @@ from . endpoint import ServiceEndpoint
from . requestor import ServiceRequestor from . requestor import ServiceRequestor
class DbpediaRequestor(ServiceRequestor): class DbpediaRequestor(ServiceRequestor):
def __init__(self, pulsar_host, timeout, auth): def __init__(self, pulsar_host, timeout, auth, pulsar_api_key=None):
super(DbpediaRequestor, self).__init__( super(DbpediaRequestor, self).__init__(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
request_queue=dbpedia_lookup_request_queue, request_queue=dbpedia_lookup_request_queue,
response_queue=dbpedia_lookup_response_queue, response_queue=dbpedia_lookup_response_queue,
request_schema=LookupRequest, request_schema=LookupRequest,

View file

@ -8,10 +8,11 @@ from . sender import ServiceSender
from . serialize import to_subgraph from . serialize import to_subgraph
class DocumentLoadSender(ServiceSender): class DocumentLoadSender(ServiceSender):
def __init__(self, pulsar_host): def __init__(self, pulsar_host, pulsar_api_key=None):
super(DocumentLoadSender, self).__init__( super(DocumentLoadSender, self).__init__(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
request_queue=document_ingest_queue, request_queue=document_ingest_queue,
request_schema=Document, request_schema=Document,
) )

View file

@ -7,10 +7,11 @@ from . endpoint import ServiceEndpoint
from . requestor import ServiceRequestor from . requestor import ServiceRequestor
class EmbeddingsRequestor(ServiceRequestor): class EmbeddingsRequestor(ServiceRequestor):
def __init__(self, pulsar_host, timeout, auth): def __init__(self, pulsar_host, timeout, auth, pulsar_api_key=None):
super(EmbeddingsRequestor, self).__init__( super(EmbeddingsRequestor, self).__init__(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
request_queue=embeddings_request_queue, request_queue=embeddings_request_queue,
response_queue=embeddings_response_queue, response_queue=embeddings_response_queue,
request_schema=EmbeddingsRequest, request_schema=EmbeddingsRequest,

View file

@ -7,10 +7,11 @@ from . endpoint import ServiceEndpoint
from . requestor import ServiceRequestor from . requestor import ServiceRequestor
class EncyclopediaRequestor(ServiceRequestor): class EncyclopediaRequestor(ServiceRequestor):
def __init__(self, pulsar_host, timeout, auth): def __init__(self, pulsar_host, timeout, auth, pulsar_api_key=None):
super(EncyclopediaRequestor, self).__init__( super(EncyclopediaRequestor, self).__init__(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
request_queue=encyclopedia_lookup_request_queue, request_queue=encyclopedia_lookup_request_queue,
response_queue=encyclopedia_lookup_response_queue, response_queue=encyclopedia_lookup_response_queue,
request_schema=LookupRequest, request_schema=LookupRequest,

View file

@ -15,7 +15,7 @@ from . serialize import to_subgraph, to_value
class GraphEmbeddingsLoadEndpoint(SocketEndpoint): class GraphEmbeddingsLoadEndpoint(SocketEndpoint):
def __init__( def __init__(
self, pulsar_host, auth, path="/api/v1/load/graph-embeddings", self, pulsar_host, auth, pulsar_api_key=None, path="/api/v1/load/graph-embeddings",
): ):
super(GraphEmbeddingsLoadEndpoint, self).__init__( super(GraphEmbeddingsLoadEndpoint, self).__init__(
@ -23,9 +23,11 @@ class GraphEmbeddingsLoadEndpoint(SocketEndpoint):
) )
self.pulsar_host=pulsar_host self.pulsar_host=pulsar_host
self.pulsar_api_key=pulsar_api_key
self.publisher = Publisher( self.publisher = Publisher(
self.pulsar_host, graph_embeddings_store_queue, self.pulsar_host, graph_embeddings_store_queue,
self.pulsar_api_key,
schema=JsonSchema(GraphEmbeddings) schema=JsonSchema(GraphEmbeddings)
) )

View file

@ -8,10 +8,11 @@ from . requestor import ServiceRequestor
from . serialize import serialize_value from . serialize import serialize_value
class GraphEmbeddingsQueryRequestor(ServiceRequestor): class GraphEmbeddingsQueryRequestor(ServiceRequestor):
def __init__(self, pulsar_host, timeout, auth): def __init__(self, pulsar_host, timeout, auth, pulsar_api_key=None):
super(GraphEmbeddingsQueryRequestor, self).__init__( super(GraphEmbeddingsQueryRequestor, self).__init__(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
request_queue=graph_embeddings_request_queue, request_queue=graph_embeddings_request_queue,
response_queue=graph_embeddings_response_queue, response_queue=graph_embeddings_response_queue,
request_schema=GraphEmbeddingsRequest, request_schema=GraphEmbeddingsRequest,

View file

@ -14,7 +14,7 @@ from . serialize import serialize_graph_embeddings
class GraphEmbeddingsStreamEndpoint(SocketEndpoint): class GraphEmbeddingsStreamEndpoint(SocketEndpoint):
def __init__( def __init__(
self, pulsar_host, auth, path="/api/v1/stream/graph-embeddings" self, pulsar_host, auth, path="/api/v1/stream/graph-embeddings", pulsar_api_key=None
): ):
super(GraphEmbeddingsStreamEndpoint, self).__init__( super(GraphEmbeddingsStreamEndpoint, self).__init__(
@ -22,10 +22,12 @@ class GraphEmbeddingsStreamEndpoint(SocketEndpoint):
) )
self.pulsar_host=pulsar_host self.pulsar_host=pulsar_host
self.pulsar_api_key=pulsar_api_key
self.subscriber = Subscriber( self.subscriber = Subscriber(
self.pulsar_host, graph_embeddings_store_queue, self.pulsar_host, graph_embeddings_store_queue,
"api-gateway", "api-gateway", "api-gateway", "api-gateway",
pulsar_api_key=self.pulsar_api_key,
schema=JsonSchema(GraphEmbeddings) schema=JsonSchema(GraphEmbeddings)
) )

View file

@ -7,10 +7,11 @@ from . endpoint import ServiceEndpoint
from . requestor import ServiceRequestor from . requestor import ServiceRequestor
class GraphRagRequestor(ServiceRequestor): class GraphRagRequestor(ServiceRequestor):
def __init__(self, pulsar_host, timeout, auth): def __init__(self, pulsar_host, timeout, auth, pulsar_api_key=None):
super(GraphRagRequestor, self).__init__( super(GraphRagRequestor, self).__init__(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
request_queue=graph_rag_request_queue, request_queue=graph_rag_request_queue,
response_queue=graph_rag_response_queue, response_queue=graph_rag_response_queue,
request_schema=GraphRagQuery, request_schema=GraphRagQuery,

View file

@ -7,10 +7,11 @@ from . endpoint import ServiceEndpoint
from . requestor import ServiceRequestor from . requestor import ServiceRequestor
class InternetSearchRequestor(ServiceRequestor): class InternetSearchRequestor(ServiceRequestor):
def __init__(self, pulsar_host, timeout, auth): def __init__(self, pulsar_host, timeout, auth, pulsar_api_key=None):
super(InternetSearchRequestor, self).__init__( super(InternetSearchRequestor, self).__init__(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
request_queue=internet_search_request_queue, request_queue=internet_search_request_queue,
response_queue=internet_search_response_queue, response_queue=internet_search_response_queue,
request_schema=LookupRequest, request_schema=LookupRequest,

View file

@ -21,6 +21,7 @@ class MuxEndpoint(SocketEndpoint):
self, pulsar_host, auth, self, pulsar_host, auth,
services, services,
path="/api/v1/socket", path="/api/v1/socket",
pulsar_api_key=None
): ):
super(MuxEndpoint, self).__init__( super(MuxEndpoint, self).__init__(

View file

@ -9,10 +9,11 @@ from . endpoint import ServiceEndpoint
from . requestor import ServiceRequestor from . requestor import ServiceRequestor
class PromptRequestor(ServiceRequestor): class PromptRequestor(ServiceRequestor):
def __init__(self, pulsar_host, timeout, auth): def __init__(self, pulsar_host, timeout, auth, pulsar_api_key=None):
super(PromptRequestor, self).__init__( super(PromptRequestor, self).__init__(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
request_queue=prompt_request_queue, request_queue=prompt_request_queue,
response_queue=prompt_response_queue, response_queue=prompt_response_queue,
request_schema=PromptRequest, request_schema=PromptRequest,

View file

@ -7,8 +7,9 @@ import threading
class Publisher: class Publisher:
def __init__(self, pulsar_host, topic, schema=None, max_size=10, def __init__(self, pulsar_host, topic, schema=None, max_size=10,
chunking_enabled=False): chunking_enabled=False, pulsar_api_key=None):
self.pulsar_host = pulsar_host self.pulsar_host = pulsar_host
self.pulsar_api_key = pulsar_api_key,
self.topic = topic self.topic = topic
self.schema = schema self.schema = schema
self.q = queue.Queue(maxsize=max_size) self.q = queue.Queue(maxsize=max_size)
@ -23,10 +24,16 @@ class Publisher:
while True: while True:
try: try:
client = pulsar.Client( if self.pulsar_api_key:
self.pulsar_host, client = pulsar.Client(
) self.pulsar_host,
authentication=pulsar.AuthenticationToken(self.pulsar_api_key)
)
else:
client = pulsar.Client(
self.pulsar_host,
)
producer = client.create_producer( producer = client.create_producer(
topic=self.topic, topic=self.topic,

View file

@ -19,16 +19,19 @@ class ServiceRequestor:
response_queue, response_schema, response_queue, response_schema,
subscription="api-gateway", consumer_name="api-gateway", subscription="api-gateway", consumer_name="api-gateway",
timeout=600, timeout=600,
pulsar_api_key=None,
): ):
self.pub = Publisher( self.pub = Publisher(
pulsar_host, request_queue, pulsar_host, request_queue,
schema=JsonSchema(request_schema) pulsar_api_key,
schema=JsonSchema(request_schema),
) )
self.sub = Subscriber( self.sub = Subscriber(
pulsar_host, response_queue, pulsar_host, response_queue,
subscription, consumer_name, subscription, consumer_name,
pulsar_api_key,
JsonSchema(response_schema) JsonSchema(response_schema)
) )

View file

@ -17,10 +17,12 @@ class ServiceSender:
self, self,
pulsar_host, pulsar_host,
request_queue, request_schema, request_queue, request_schema,
pulsar_api_key=None,
): ):
self.pub = Publisher( self.pub = Publisher(
pulsar_host, request_queue, pulsar_host, request_queue,
pulsar_api_key,
schema=JsonSchema(request_schema) schema=JsonSchema(request_schema)
) )

View file

@ -53,6 +53,7 @@ logger = logging.getLogger("api")
logger.setLevel(logging.INFO) logger.setLevel(logging.INFO)
default_pulsar_host = os.getenv("PULSAR_HOST", "pulsar://pulsar:6650") default_pulsar_host = os.getenv("PULSAR_HOST", "pulsar://pulsar:6650")
default_pulsar_api_key = os.getenv("PULSAR_API_KEY", None)
default_timeout = 600 default_timeout = 600
default_port = 8088 default_port = 8088
default_api_token = os.getenv("GATEWAY_SECRET", "") default_api_token = os.getenv("GATEWAY_SECRET", "")
@ -69,6 +70,7 @@ class Api:
self.port = int(config.get("port", default_port)) self.port = int(config.get("port", default_port))
self.timeout = int(config.get("timeout", default_timeout)) self.timeout = int(config.get("timeout", default_timeout))
self.pulsar_host = config.get("pulsar_host", default_pulsar_host) self.pulsar_host = config.get("pulsar_host", default_pulsar_host)
self.pulsar_api_key = config.get("pulsar_api_key", default_pulsar_api_key)
api_token = config.get("api_token", default_api_token) api_token = config.get("api_token", default_api_token)
@ -81,49 +83,49 @@ class Api:
self.services = { self.services = {
"text-completion": TextCompletionRequestor( "text-completion": TextCompletionRequestor(
pulsar_host=self.pulsar_host, timeout=self.timeout, pulsar_host=self.pulsar_host, timeout=self.timeout,
auth = self.auth, auth = self.auth, pulsar_api_key=self.pulsar_api_key,
), ),
"prompt": PromptRequestor( "prompt": PromptRequestor(
pulsar_host=self.pulsar_host, timeout=self.timeout, pulsar_host=self.pulsar_host, timeout=self.timeout,
auth = self.auth, auth = self.auth, pulsar_api_key=self.pulsar_api_key,
), ),
"graph-rag": GraphRagRequestor( "graph-rag": GraphRagRequestor(
pulsar_host=self.pulsar_host, timeout=self.timeout, pulsar_host=self.pulsar_host, timeout=self.timeout,
auth = self.auth, auth = self.auth, pulsar_api_key=self.pulsar_api_key,
), ),
"triples-query": TriplesQueryRequestor( "triples-query": TriplesQueryRequestor(
pulsar_host=self.pulsar_host, timeout=self.timeout, pulsar_host=self.pulsar_host, timeout=self.timeout,
auth = self.auth, auth = self.auth, pulsar_api_key=self.pulsar_api_key,
), ),
"graph-embeddings-query": GraphEmbeddingsQueryRequestor( "graph-embeddings-query": GraphEmbeddingsQueryRequestor(
pulsar_host=self.pulsar_host, timeout=self.timeout, pulsar_host=self.pulsar_host, timeout=self.timeout,
auth = self.auth, auth = self.auth, pulsar_api_key=self.pulsar_api_key,
), ),
"embeddings": EmbeddingsRequestor( "embeddings": EmbeddingsRequestor(
pulsar_host=self.pulsar_host, timeout=self.timeout, pulsar_host=self.pulsar_host, timeout=self.timeout,
auth = self.auth, auth = self.auth, pulsar_api_key=self.pulsar_api_key,
), ),
"agent": AgentRequestor( "agent": AgentRequestor(
pulsar_host=self.pulsar_host, timeout=self.timeout, pulsar_host=self.pulsar_host, timeout=self.timeout,
auth = self.auth, auth = self.auth, pulsar_api_key=self.pulsar_api_key,
), ),
"encyclopedia": EncyclopediaRequestor( "encyclopedia": EncyclopediaRequestor(
pulsar_host=self.pulsar_host, timeout=self.timeout, pulsar_host=self.pulsar_host, timeout=self.timeout,
auth = self.auth, auth = self.auth, pulsar_api_key=self.pulsar_api_key,
), ),
"dbpedia": DbpediaRequestor( "dbpedia": DbpediaRequestor(
pulsar_host=self.pulsar_host, timeout=self.timeout, pulsar_host=self.pulsar_host, timeout=self.timeout,
auth = self.auth, auth = self.auth, pulsar_api_key=self.pulsar_api_key,
), ),
"internet-search": InternetSearchRequestor( "internet-search": InternetSearchRequestor(
pulsar_host=self.pulsar_host, timeout=self.timeout, pulsar_host=self.pulsar_host, timeout=self.timeout,
auth = self.auth, auth = self.auth, pulsar_api_key=self.pulsar_api_key,
), ),
"document-load": DocumentLoadSender( "document-load": DocumentLoadSender(
pulsar_host=self.pulsar_host, pulsar_host=self.pulsar_host, pulsar_api_key=self.pulsar_api_key,
), ),
"text-load": TextLoadSender( "text-load": TextLoadSender(
pulsar_host=self.pulsar_host, pulsar_host=self.pulsar_host, pulsar_api_key=self.pulsar_api_key,
), ),
} }
@ -179,24 +181,29 @@ class Api:
), ),
TriplesStreamEndpoint( TriplesStreamEndpoint(
pulsar_host=self.pulsar_host, pulsar_host=self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
auth = self.auth, auth = self.auth,
), ),
GraphEmbeddingsStreamEndpoint( GraphEmbeddingsStreamEndpoint(
pulsar_host=self.pulsar_host, pulsar_host=self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
auth = self.auth, auth = self.auth,
), ),
TriplesLoadEndpoint( TriplesLoadEndpoint(
pulsar_host=self.pulsar_host, pulsar_host=self.pulsar_host,
auth = self.auth, auth = self.auth,
pulsar_api_key=self.pulsar_api_key,
), ),
GraphEmbeddingsLoadEndpoint( GraphEmbeddingsLoadEndpoint(
pulsar_host=self.pulsar_host, pulsar_host=self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
auth = self.auth, auth = self.auth,
), ),
MuxEndpoint( MuxEndpoint(
pulsar_host=self.pulsar_host, pulsar_host=self.pulsar_host,
auth = self.auth, auth = self.auth,
services = self.services, services = self.services,
pulsar_api_key=self.pulsar_api_key,
), ),
] ]
@ -225,6 +232,12 @@ def run():
default=default_pulsar_host, default=default_pulsar_host,
help=f'Pulsar host (default: {default_pulsar_host})', help=f'Pulsar host (default: {default_pulsar_host})',
) )
parser.add_argument(
'--pulsar-api-key',
default=default_pulsar_api_key,
help=f'Pulsar API key',
)
parser.add_argument( parser.add_argument(
'--port', '--port',

View file

@ -6,9 +6,10 @@ import time
class Subscriber: class Subscriber:
def __init__(self, pulsar_host, topic, subscription, consumer_name, def __init__(self, pulsar_host, topic, subscription, consumer_name, pulsar_api_key=None,
schema=None, max_size=100): schema=None, max_size=100):
self.pulsar_host = pulsar_host self.pulsar_host = pulsar_host
self.pulsar_api_key = pulsar_api_key
self.topic = topic self.topic = topic
self.subscription = subscription self.subscription = subscription
self.consumer_name = consumer_name self.consumer_name = consumer_name
@ -28,9 +29,16 @@ class Subscriber:
try: try:
client = pulsar.Client( if self.pulsar_api_key:
auth = pulsar.AuthenticationToken(self.pulsar_api_key)
client = pulsar.Client(
self.pulsar_host, self.pulsar_host,
) authentication=auth,
)
else:
client = pulsar.Client(
self.pulsar_host,
)
consumer = client.subscribe( consumer = client.subscribe(
topic=self.topic, topic=self.topic,

View file

@ -7,10 +7,11 @@ from . endpoint import ServiceEndpoint
from . requestor import ServiceRequestor from . requestor import ServiceRequestor
class TextCompletionRequestor(ServiceRequestor): class TextCompletionRequestor(ServiceRequestor):
def __init__(self, pulsar_host, timeout, auth): def __init__(self, pulsar_host, timeout, auth, pulsar_api_key=None):
super(TextCompletionRequestor, self).__init__( super(TextCompletionRequestor, self).__init__(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
request_queue=text_completion_request_queue, request_queue=text_completion_request_queue,
response_queue=text_completion_response_queue, response_queue=text_completion_response_queue,
request_schema=TextCompletionRequest, request_schema=TextCompletionRequest,

View file

@ -8,10 +8,11 @@ from . sender import ServiceSender
from . serialize import to_subgraph from . serialize import to_subgraph
class TextLoadSender(ServiceSender): class TextLoadSender(ServiceSender):
def __init__(self, pulsar_host): def __init__(self, pulsar_host, pulsar_api_key=None):
super(TextLoadSender, self).__init__( super(TextLoadSender, self).__init__(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
request_queue=text_ingest_queue, request_queue=text_ingest_queue,
request_schema=TextDocument, request_schema=TextDocument,
) )

View file

@ -14,17 +14,19 @@ from . serialize import to_subgraph
class TriplesLoadEndpoint(SocketEndpoint): class TriplesLoadEndpoint(SocketEndpoint):
def __init__(self, pulsar_host, auth, path="/api/v1/load/triples"): def __init__(self, pulsar_host, auth, path="/api/v1/load/triples", pulsar_api_key=None):
super(TriplesLoadEndpoint, self).__init__( super(TriplesLoadEndpoint, self).__init__(
endpoint_path=path, auth=auth, endpoint_path=path, auth=auth,
) )
self.pulsar_host=pulsar_host self.pulsar_host=pulsar_host
self.pulsar_api_key=pulsar_api_key
self.publisher = Publisher( self.publisher = Publisher(
self.pulsar_host, triples_store_queue, self.pulsar_host, triples_store_queue,
schema=JsonSchema(Triples) schema=JsonSchema(Triples),
pulsar_api_key=self.pulsar_api_key
) )
async def start(self): async def start(self):

View file

@ -8,10 +8,11 @@ from . requestor import ServiceRequestor
from . serialize import to_value, serialize_subgraph from . serialize import to_value, serialize_subgraph
class TriplesQueryRequestor(ServiceRequestor): class TriplesQueryRequestor(ServiceRequestor):
def __init__(self, pulsar_host, timeout, auth): def __init__(self, pulsar_host, timeout, auth, pulsar_api_key=None):
super(TriplesQueryRequestor, self).__init__( super(TriplesQueryRequestor, self).__init__(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=pulsar_api_key,
request_queue=triples_request_queue, request_queue=triples_request_queue,
response_queue=triples_response_queue, response_queue=triples_response_queue,
request_schema=TriplesQueryRequest, request_schema=TriplesQueryRequest,

View file

@ -13,17 +13,19 @@ from . serialize import serialize_triples
class TriplesStreamEndpoint(SocketEndpoint): class TriplesStreamEndpoint(SocketEndpoint):
def __init__(self, pulsar_host, auth, path="/api/v1/stream/triples"): def __init__(self, pulsar_host, auth, path="/api/v1/stream/triples", pulsar_api_key=None):
super(TriplesStreamEndpoint, self).__init__( super(TriplesStreamEndpoint, self).__init__(
endpoint_path=path, auth=auth, endpoint_path=path, auth=auth,
) )
self.pulsar_host=pulsar_host self.pulsar_host=pulsar_host
self.pulsar_api_key=pulsar_api_key
self.subscriber = Subscriber( self.subscriber = Subscriber(
self.pulsar_host, triples_store_queue, self.pulsar_host, triples_store_queue,
"api-gateway", "api-gateway", "api-gateway", "api-gateway",
pulsar_api_key=self.pulsar_api_key,
schema=JsonSchema(Triples) schema=JsonSchema(Triples)
) )

View file

@ -161,6 +161,7 @@ class GraphRag:
def __init__( def __init__(
self, self,
pulsar_host="pulsar://pulsar:6650", pulsar_host="pulsar://pulsar:6650",
pulsar_api_key=None,
pr_request_queue=None, pr_request_queue=None,
pr_response_queue=None, pr_response_queue=None,
emb_request_queue=None, emb_request_queue=None,
@ -207,6 +208,7 @@ class GraphRag:
self.ge_client = GraphEmbeddingsClient( self.ge_client = GraphEmbeddingsClient(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=-pulsar_api_key,
subscriber=module + "-ge", subscriber=module + "-ge",
input_queue=ge_request_queue, input_queue=ge_request_queue,
output_queue=ge_response_queue, output_queue=ge_response_queue,
@ -214,6 +216,7 @@ class GraphRag:
self.triples_client = TriplesQueryClient( self.triples_client = TriplesQueryClient(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=-pulsar_api_key,
subscriber=module + "-tpl", subscriber=module + "-tpl",
input_queue=tpl_request_queue, input_queue=tpl_request_queue,
output_queue=tpl_response_queue output_queue=tpl_response_queue
@ -221,6 +224,7 @@ class GraphRag:
self.embeddings = EmbeddingsClient( self.embeddings = EmbeddingsClient(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=-pulsar_api_key,
input_queue=emb_request_queue, input_queue=emb_request_queue,
output_queue=emb_response_queue, output_queue=emb_response_queue,
subscriber=module + "-emb", subscriber=module + "-emb",
@ -234,6 +238,7 @@ class GraphRag:
self.prompt = PromptClient( self.prompt = PromptClient(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
pulsar_api_key=-pulsar_api_key,
input_queue=pr_request_queue, input_queue=pr_request_queue,
output_queue=pr_response_queue, output_queue=pr_response_queue,
subscriber=module + "-prompt", subscriber=module + "-prompt",

View file

@ -63,7 +63,8 @@ class Processor(ConsumerProducer):
subscriber=subscriber, subscriber=subscriber,
input_queue=tc_request_queue, input_queue=tc_request_queue,
output_queue=tc_response_queue, output_queue=tc_response_queue,
pulsar_host = self.pulsar_host pulsar_host = self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
) )
def parse_json(self, text): def parse_json(self, text):

View file

@ -136,7 +136,8 @@ class Processor(ConsumerProducer):
subscriber=subscriber, subscriber=subscriber,
input_queue=tc_request_queue, input_queue=tc_request_queue,
output_queue=tc_response_queue, output_queue=tc_response_queue,
pulsar_host = self.pulsar_host pulsar_host = self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
) )
# System prompt hack # System prompt hack

View file

@ -49,11 +49,12 @@ class Processing:
pulsar_host, pulsar_host,
log_level, log_level,
file, file,
pulsar_api_key=None,
): ):
self.pulsar_host = pulsar_host self.pulsar_host = pulsar_host
self.log_level = log_level self.log_level = log_level
self.file = file self.file = file
self.pulsar_api_key = pulsar_api_key
self.defs = load(open(file, "r"), Loader=Loader) self.defs = load(open(file, "r"), Loader=Loader)
def run(self): def run(self):
@ -125,12 +126,19 @@ def run():
) )
default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://pulsar:6650') default_pulsar_host = os.getenv("PULSAR_HOST", 'pulsar://pulsar:6650')
default_pulsar_api_key = os.getenv("PULSAR_API_KEY", None)
parser.add_argument( parser.add_argument(
'-p', '--pulsar-host', '-p', '--pulsar-host',
default=default_pulsar_host, default=default_pulsar_host,
help=f'Pulsar host (default: {default_pulsar_host})', help=f'Pulsar host (default: {default_pulsar_host})',
) )
parser.add_argument(
'--pulsar-api-key',
default=default_pulsar_api_key,
help=f'Pulsar API key',
)
parser.add_argument( parser.add_argument(
'-l', '--log-level', '-l', '--log-level',

View file

@ -68,6 +68,7 @@ class Processor(ConsumerProducer):
self.rag = DocumentRag( self.rag = DocumentRag(
pulsar_host=self.pulsar_host, pulsar_host=self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
pr_request_queue=pr_request_queue, pr_request_queue=pr_request_queue,
pr_response_queue=pr_response_queue, pr_response_queue=pr_response_queue,
emb_request_queue=emb_request_queue, emb_request_queue=emb_request_queue,

View file

@ -82,6 +82,7 @@ class Processor(ConsumerProducer):
self.rag = GraphRag( self.rag = GraphRag(
pulsar_host=self.pulsar_host, pulsar_host=self.pulsar_host,
pulsar_api_key=self.pulsar_api_key,
pr_request_queue=pr_request_queue, pr_request_queue=pr_request_queue,
pr_response_queue=pr_response_queue, pr_response_queue=pr_response_queue,
emb_request_queue=emb_request_queue, emb_request_queue=emb_request_queue,