From 8ab0a9ace6e1b47144ab74bcd0e9036bb72832af Mon Sep 17 00:00:00 2001 From: Cyber MacGeddon Date: Fri, 25 Apr 2025 16:21:46 +0100 Subject: [PATCH] Refactoring some LLM handlers --- .../trustgraph/base/llm_service.py | 5 + .../model/text_completion/azure/llm.py | 87 +++--------- .../model/text_completion/azure_openai/llm.py | 123 +++++------------ .../model/text_completion/claude/llm.py | 124 ++++-------------- .../model/text_completion/vertexai/llm.py | 11 +- 5 files changed, 91 insertions(+), 259 deletions(-) diff --git a/trustgraph-base/trustgraph/base/llm_service.py b/trustgraph-base/trustgraph/base/llm_service.py index 39323db7..c79b819b 100644 --- a/trustgraph-base/trustgraph/base/llm_service.py +++ b/trustgraph-base/trustgraph/base/llm_service.py @@ -13,6 +13,11 @@ from .. base import FlowProcessor, ConsumerSpec, ProducerSpec default_ident = "text-completion" class LlmResult: + def __init__(self, text=None, in_token=None, out_token=None, model=None): + self.text = text + self.in_token = in_token + self.out_token = out_token + self.model = model __slots__ = ["text", "in_token", "out_token", "model"] class LlmService(FlowProcessor): diff --git a/trustgraph-flow/trustgraph/model/text_completion/azure/llm.py b/trustgraph-flow/trustgraph/model/text_completion/azure/llm.py index 79118cc8..70b07606 100755 --- a/trustgraph-flow/trustgraph/model/text_completion/azure/llm.py +++ b/trustgraph-flow/trustgraph/model/text_completion/azure/llm.py @@ -9,31 +9,21 @@ import json from prometheus_client import Histogram import os -from .... schema import TextCompletionRequest, TextCompletionResponse, Error -from .... schema import text_completion_request_queue -from .... schema import text_completion_response_queue -from .... log_level import LogLevel -from .... base import ConsumerProducer from .... exceptions import TooManyRequests +from .... base import LlmService, LlmResult -module = "text-completion" +default_ident = "text-completion" -default_input_queue = text_completion_request_queue -default_output_queue = text_completion_response_queue -default_subscriber = module default_temperature = 0.0 default_max_output = 4192 default_model = "AzureAI" default_endpoint = os.getenv("AZURE_ENDPOINT") default_token = os.getenv("AZURE_TOKEN") -class Processor(ConsumerProducer): +class Processor(LlmService): 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) endpoint = params.get("endpoint", default_endpoint) token = params.get("token", default_token) temperature = params.get("temperature", default_temperature) @@ -48,30 +38,13 @@ class Processor(ConsumerProducer): super(Processor, self).__init__( **params | { - "input_queue": input_queue, - "output_queue": output_queue, - "subscriber": subscriber, - "input_schema": TextCompletionRequest, - "output_schema": TextCompletionResponse, + "endpoint": endpoint, "temperature": temperature, "max_output": max_output, "model": model, } ) - if not hasattr(__class__, "text_completion_metric"): - __class__.text_completion_metric = Histogram( - 'text_completion_duration', - 'Text completion duration (seconds)', - buckets=[ - 0.25, 0.5, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, - 8.0, 9.0, 10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0, - 17.0, 18.0, 19.0, 20.0, 21.0, 22.0, 23.0, 24.0, 25.0, - 30.0, 35.0, 40.0, 45.0, 50.0, 60.0, 80.0, 100.0, - 120.0 - ] - ) - self.endpoint = endpoint self.token = token self.temperature = temperature @@ -123,25 +96,16 @@ class Processor(ConsumerProducer): return result - async def handle(self, msg): - - v = msg.value() - - # Sender-produced ID - - id = msg.properties()["id"] - - print(f"Handling prompt {id}...", flush=True) + async def generate_content(self, system, prompt): try: prompt = self.build_prompt( - v.system, - v.prompt + system, + prompt ) - with __class__.text_completion_metric.time(): - response = self.call_llm(prompt) + response = self.call_llm(prompt) resp = response['choices'][0]['message']['content'] inputtokens = response['usage']['prompt_tokens'] @@ -153,8 +117,14 @@ class Processor(ConsumerProducer): print("Send response...", flush=True) - r = TextCompletionResponse(response=resp, error=None, in_token=inputtokens, out_token=outputtokens, model=self.model) - await self.send(r, properties={"id": id}) + resp = LlmResult( + text = resp, + in_token = inputtokens, + out_token = outputtokens, + model = self.model + ) + + return resp except TooManyRequests: @@ -168,33 +138,14 @@ class Processor(ConsumerProducer): # Apart from rate limits, treat all exceptions as unrecoverable print(f"Exception: {e}") - - print("Send error response...", flush=True) - - r = TextCompletionResponse( - error=Error( - type = "llm-error", - message = str(e), - ), - response=None, - in_token=None, - out_token=None, - model=None, - ) - - await self.send(r, properties={"id": id}) - - self.consumer.acknowledge(msg) + raise e print("Done.", flush=True) @staticmethod def add_args(parser): - ConsumerProducer.add_args( - parser, default_input_queue, default_subscriber, - default_output_queue, - ) + LlmService.add_args(parser) parser.add_argument( '-e', '--endpoint', @@ -224,4 +175,4 @@ class Processor(ConsumerProducer): def run(): - Processor.launch(module, __doc__) + Processor.launch(default_ident, __doc__) diff --git a/trustgraph-flow/trustgraph/model/text_completion/azure_openai/llm.py b/trustgraph-flow/trustgraph/model/text_completion/azure_openai/llm.py index 734b20c5..c5dd097c 100755 --- a/trustgraph-flow/trustgraph/model/text_completion/azure_openai/llm.py +++ b/trustgraph-flow/trustgraph/model/text_completion/azure_openai/llm.py @@ -9,18 +9,11 @@ from prometheus_client import Histogram from openai import AzureOpenAI, RateLimitError import os -from .... schema import TextCompletionRequest, TextCompletionResponse, Error -from .... schema import text_completion_request_queue -from .... schema import text_completion_response_queue -from .... log_level import LogLevel -from .... base import ConsumerProducer from .... exceptions import TooManyRequests +from .... base import LlmService, LlmResult -module = "text-completion" +default_ident = "text-completion" -default_input_queue = text_completion_request_queue -default_output_queue = text_completion_response_queue -default_subscriber = module default_temperature = 0.0 default_max_output = 4192 default_api = "2024-12-01-preview" @@ -28,13 +21,10 @@ default_endpoint = os.getenv("AZURE_ENDPOINT", None) default_token = os.getenv("AZURE_TOKEN", None) default_model = os.getenv("AZURE_MODEL", None) -class Processor(ConsumerProducer): +class Processor(LlmService): 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) temperature = params.get("temperature", default_temperature) max_output = params.get("max_output", default_max_output) @@ -51,11 +41,6 @@ class Processor(ConsumerProducer): super(Processor, self).__init__( **params | { - "input_queue": input_queue, - "output_queue": output_queue, - "subscriber": subscriber, - "input_schema": TextCompletionRequest, - "output_schema": TextCompletionResponse, "temperature": temperature, "max_output": max_output, "model": model, @@ -63,19 +48,6 @@ class Processor(ConsumerProducer): } ) - if not hasattr(__class__, "text_completion_metric"): - __class__.text_completion_metric = Histogram( - 'text_completion_duration', - 'Text completion duration (seconds)', - buckets=[ - 0.25, 0.5, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, - 8.0, 9.0, 10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0, - 17.0, 18.0, 19.0, 20.0, 21.0, 22.0, 23.0, 24.0, 25.0, - 30.0, 35.0, 40.0, 45.0, 50.0, 60.0, 80.0, 100.0, - 120.0 - ] - ) - self.temperature = temperature self.max_output = max_output self.model = model @@ -84,41 +56,31 @@ class Processor(ConsumerProducer): api_key=token, api_version=api, azure_endpoint = endpoint, - ) + ) - async def handle(self, msg): - - v = msg.value() - - # Sender-produced ID - - id = msg.properties()["id"] - - print(f"Handling prompt {id}...", flush=True) - - prompt = v.system + "\n\n" + v.prompt + async def generate_content(self, system, prompt): + prompt = system + "\n\n" + prompt try: - with __class__.text_completion_metric.time(): - resp = self.openai.chat.completions.create( - model=self.model, - messages=[ - { - "role": "user", - "content": [ - { - "type": "text", - "text": prompt - } - ] - } - ], - temperature=self.temperature, - max_tokens=self.max_output, - top_p=1, - ) + resp = self.openai.chat.completions.create( + model=self.model, + messages=[ + { + "role": "user", + "content": [ + { + "type": "text", + "text": prompt + } + ] + } + ], + temperature=self.temperature, + max_tokens=self.max_output, + top_p=1, + ) inputtokens = resp.usage.prompt_tokens outputtokens = resp.usage.completion_tokens @@ -127,15 +89,14 @@ 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 + r = LlmResult( + text = resp.choices[0].message.content, + in_token = inputtokens, + out_token = outputtokens, + model = self.model ) - await self.send(r, properties={"id": id}) + return r except RateLimitError: @@ -147,35 +108,15 @@ class Processor(ConsumerProducer): except Exception as e: # Apart from rate limits, treat all exceptions as unrecoverable - print(f"Exception: {e}") - - print("Send error response...", flush=True) - - r = TextCompletionResponse( - error=Error( - type = "llm-error", - message = str(e), - ), - response=None, - in_token=None, - out_token=None, - model=None, - ) - - await self.send(r, properties={"id": id}) - - self.consumer.acknowledge(msg) + raise e print("Done.", flush=True) @staticmethod def add_args(parser): - ConsumerProducer.add_args( - parser, default_input_queue, default_subscriber, - default_output_queue, - ) + LlmService.add_args(parser) parser.add_argument( '-e', '--endpoint', @@ -217,4 +158,4 @@ class Processor(ConsumerProducer): def run(): - Processor.launch(module, __doc__) + Processor.launch(default_ident, __doc__) diff --git a/trustgraph-flow/trustgraph/model/text_completion/claude/llm.py b/trustgraph-flow/trustgraph/model/text_completion/claude/llm.py index f60b70d7..eb303ac2 100755 --- a/trustgraph-flow/trustgraph/model/text_completion/claude/llm.py +++ b/trustgraph-flow/trustgraph/model/text_completion/claude/llm.py @@ -8,30 +8,20 @@ import anthropic from prometheus_client import Histogram import os -from .... schema import TextCompletionRequest, TextCompletionResponse, Error -from .... schema import text_completion_request_queue -from .... schema import text_completion_response_queue -from .... log_level import LogLevel -from .... base import ConsumerProducer from .... exceptions import TooManyRequests +from .... base import LlmService, LlmResult -module = "text-completion" +default_ident = "text-completion" -default_input_queue = text_completion_request_queue -default_output_queue = text_completion_response_queue -default_subscriber = module default_model = 'claude-3-5-sonnet-20240620' default_temperature = 0.0 default_max_output = 8192 default_api_key = os.getenv("CLAUDE_KEY") -class Processor(ConsumerProducer): +class Processor(LlmService): 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) api_key = params.get("api_key", default_api_key) temperature = params.get("temperature", default_temperature) @@ -42,30 +32,12 @@ class Processor(ConsumerProducer): super(Processor, self).__init__( **params | { - "input_queue": input_queue, - "output_queue": output_queue, - "subscriber": subscriber, - "input_schema": TextCompletionRequest, - "output_schema": TextCompletionResponse, "model": model, "temperature": temperature, "max_output": max_output, } ) - if not hasattr(__class__, "text_completion_metric"): - __class__.text_completion_metric = Histogram( - 'text_completion_duration', - 'Text completion duration (seconds)', - buckets=[ - 0.25, 0.5, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, - 8.0, 9.0, 10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0, - 17.0, 18.0, 19.0, 20.0, 21.0, 22.0, 23.0, 24.0, 25.0, - 30.0, 35.0, 40.0, 45.0, 50.0, 60.0, 80.0, 100.0, - 120.0 - ] - ) - self.model = model self.claude = anthropic.Anthropic(api_key=api_key) self.temperature = temperature @@ -73,39 +45,27 @@ class Processor(ConsumerProducer): print("Initialised", flush=True) - async def handle(self, msg): - - v = msg.value() - - # Sender-produced ID - - id = msg.properties()["id"] - - print(f"Handling prompt {id}...", flush=True) - - prompt = v.prompt + async def generate_content(self, system, prompt): try: - with __class__.text_completion_metric.time(): - - response = message = self.claude.messages.create( - model=self.model, - max_tokens=self.max_output, - temperature=self.temperature, - system = v.system, - messages=[ - { - "role": "user", - "content": [ - { - "type": "text", - "text": prompt - } - ] - } - ] - ) + response = message = self.claude.messages.create( + model=self.model, + max_tokens=self.max_output, + temperature=self.temperature, + system = system, + messages=[ + { + "role": "user", + "content": [ + { + "type": "text", + "text": prompt + } + ] + } + ] + ) resp = response.content[0].text inputtokens = response.usage.input_tokens @@ -114,17 +74,12 @@ class Processor(ConsumerProducer): print(f"Input Tokens: {inputtokens}", flush=True) 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 + resp = LlmResult( + text = resp, + in_token = inputtokens, + out_token = outputtokens, + model = self.model ) - self.send(r, properties={"id": id}) - - print("Done.", flush=True) except anthropic.RateLimitError: @@ -136,31 +91,12 @@ class Processor(ConsumerProducer): # Apart from rate limits, treat all exceptions as unrecoverable print(f"Exception: {e}") - - print("Send error response...", flush=True) - - r = TextCompletionResponse( - error=Error( - type = "llm-error", - message = str(e), - ), - response=None, - in_token=None, - out_token=None, - model=None, - ) - - await self.send(r, properties={"id": id}) - - self.consumer.acknowledge(msg) + raise e @staticmethod def add_args(parser): - ConsumerProducer.add_args( - parser, default_input_queue, default_subscriber, - default_output_queue, - ) + LlmService.add_args(parser) parser.add_argument( '-m', '--model', @@ -189,7 +125,5 @@ class Processor(ConsumerProducer): ) def run(): - - Processor.launch(module, __doc__) - + Processor.launch(default_ident, __doc__) diff --git a/trustgraph-vertexai/trustgraph/model/text_completion/vertexai/llm.py b/trustgraph-vertexai/trustgraph/model/text_completion/vertexai/llm.py index 3594b76d..854be961 100755 --- a/trustgraph-vertexai/trustgraph/model/text_completion/vertexai/llm.py +++ b/trustgraph-vertexai/trustgraph/model/text_completion/vertexai/llm.py @@ -105,11 +105,12 @@ class Processor(LlmService): safety_settings=self.safety_settings ) - resp = LlmResult() - resp.text = response.text - resp.in_token = response.usage_metadata.prompt_token_count - resp.out_token = response.usage_metadata.candidates_token_count - resp.model = self.model + resp = LlmResult( + text = response.text, + in_token = response.usage_metadata.prompt_token_count, + out_token = response.usage_metadata.candidates_token_count, + model = self.model + ) print(f"Input Tokens: {resp.in_token}", flush=True) print(f"Output Tokens: {resp.out_token}", flush=True)