mirror of
https://github.com/trustgraph-ai/trustgraph.git
synced 2026-07-24 04:31:02 +02:00
Phase 1 & 2 of streaming, covers some VertexAI prototyping
This commit is contained in:
parent
99959d34c9
commit
cab6e95d80
9 changed files with 417 additions and 50 deletions
|
|
@ -28,6 +28,19 @@ class LlmResult:
|
||||||
self.model = model
|
self.model = model
|
||||||
__slots__ = ["text", "in_token", "out_token", "model"]
|
__slots__ = ["text", "in_token", "out_token", "model"]
|
||||||
|
|
||||||
|
class LlmChunk:
|
||||||
|
"""Represents a streaming chunk from an LLM"""
|
||||||
|
def __init__(
|
||||||
|
self, text = None, in_token = None, out_token = None,
|
||||||
|
model = None, is_final = False,
|
||||||
|
):
|
||||||
|
self.text = text
|
||||||
|
self.in_token = in_token
|
||||||
|
self.out_token = out_token
|
||||||
|
self.model = model
|
||||||
|
self.is_final = is_final
|
||||||
|
__slots__ = ["text", "in_token", "out_token", "model", "is_final"]
|
||||||
|
|
||||||
class LlmService(FlowProcessor):
|
class LlmService(FlowProcessor):
|
||||||
|
|
||||||
def __init__(self, **params):
|
def __init__(self, **params):
|
||||||
|
|
@ -99,16 +112,57 @@ class LlmService(FlowProcessor):
|
||||||
|
|
||||||
id = msg.properties()["id"]
|
id = msg.properties()["id"]
|
||||||
|
|
||||||
with __class__.text_completion_metric.labels(
|
model = flow("model")
|
||||||
id=self.id,
|
temperature = flow("temperature")
|
||||||
flow=f"{flow.name}-{consumer.name}",
|
|
||||||
).time():
|
|
||||||
|
|
||||||
model = flow("model")
|
# Check if streaming is requested and supported
|
||||||
temperature = flow("temperature")
|
streaming = getattr(request, 'streaming', False)
|
||||||
|
|
||||||
response = await self.generate_content(
|
if streaming and self.supports_streaming():
|
||||||
request.system, request.prompt, model, temperature
|
|
||||||
|
# Streaming mode
|
||||||
|
with __class__.text_completion_metric.labels(
|
||||||
|
id=self.id,
|
||||||
|
flow=f"{flow.name}-{consumer.name}",
|
||||||
|
).time():
|
||||||
|
|
||||||
|
async for chunk in self.generate_content_stream(
|
||||||
|
request.system, request.prompt, model, temperature
|
||||||
|
):
|
||||||
|
await flow("response").send(
|
||||||
|
TextCompletionResponse(
|
||||||
|
error=None,
|
||||||
|
response=chunk.text,
|
||||||
|
in_token=chunk.in_token,
|
||||||
|
out_token=chunk.out_token,
|
||||||
|
model=chunk.model,
|
||||||
|
end_of_stream=chunk.is_final
|
||||||
|
),
|
||||||
|
properties={"id": id}
|
||||||
|
)
|
||||||
|
|
||||||
|
else:
|
||||||
|
|
||||||
|
# Non-streaming mode (original behavior)
|
||||||
|
with __class__.text_completion_metric.labels(
|
||||||
|
id=self.id,
|
||||||
|
flow=f"{flow.name}-{consumer.name}",
|
||||||
|
).time():
|
||||||
|
|
||||||
|
response = await self.generate_content(
|
||||||
|
request.system, request.prompt, model, temperature
|
||||||
|
)
|
||||||
|
|
||||||
|
await flow("response").send(
|
||||||
|
TextCompletionResponse(
|
||||||
|
error=None,
|
||||||
|
response=response.text,
|
||||||
|
in_token=response.in_token,
|
||||||
|
out_token=response.out_token,
|
||||||
|
model=response.model,
|
||||||
|
end_of_stream=True
|
||||||
|
),
|
||||||
|
properties={"id": id}
|
||||||
)
|
)
|
||||||
|
|
||||||
__class__.text_completion_model_metric.labels(
|
__class__.text_completion_model_metric.labels(
|
||||||
|
|
@ -119,17 +173,6 @@ class LlmService(FlowProcessor):
|
||||||
"temperature": str(temperature) if temperature is not None else "",
|
"temperature": str(temperature) if temperature is not None else "",
|
||||||
})
|
})
|
||||||
|
|
||||||
await flow("response").send(
|
|
||||||
TextCompletionResponse(
|
|
||||||
error=None,
|
|
||||||
response=response.text,
|
|
||||||
in_token=response.in_token,
|
|
||||||
out_token=response.out_token,
|
|
||||||
model=response.model
|
|
||||||
),
|
|
||||||
properties={"id": id}
|
|
||||||
)
|
|
||||||
|
|
||||||
except TooManyRequests as e:
|
except TooManyRequests as e:
|
||||||
raise e
|
raise e
|
||||||
|
|
||||||
|
|
@ -151,10 +194,26 @@ class LlmService(FlowProcessor):
|
||||||
in_token=None,
|
in_token=None,
|
||||||
out_token=None,
|
out_token=None,
|
||||||
model=None,
|
model=None,
|
||||||
|
end_of_stream=True
|
||||||
),
|
),
|
||||||
properties={"id": id}
|
properties={"id": id}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def supports_streaming(self):
|
||||||
|
"""
|
||||||
|
Override in subclass to indicate streaming support.
|
||||||
|
Returns False by default.
|
||||||
|
"""
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def generate_content_stream(self, system, prompt, model=None, temperature=None):
|
||||||
|
"""
|
||||||
|
Override in subclass to implement streaming.
|
||||||
|
Should yield LlmChunk objects.
|
||||||
|
The final chunk should have is_final=True.
|
||||||
|
"""
|
||||||
|
raise NotImplementedError("Streaming not implemented for this provider")
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def add_args(parser):
|
def add_args(parser):
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -5,6 +5,7 @@ from .. schema import TextCompletionRequest, TextCompletionResponse
|
||||||
from .. schema import text_completion_request_queue
|
from .. schema import text_completion_request_queue
|
||||||
from .. schema import text_completion_response_queue
|
from .. schema import text_completion_response_queue
|
||||||
from . base import BaseClient
|
from . base import BaseClient
|
||||||
|
from .. exceptions import LlmError
|
||||||
|
|
||||||
# Ugly
|
# Ugly
|
||||||
ERROR=_pulsar.LoggerLevel.Error
|
ERROR=_pulsar.LoggerLevel.Error
|
||||||
|
|
@ -37,8 +38,68 @@ class LlmClient(BaseClient):
|
||||||
output_schema=TextCompletionResponse,
|
output_schema=TextCompletionResponse,
|
||||||
)
|
)
|
||||||
|
|
||||||
def request(self, system, prompt, timeout=300):
|
def request(self, system, prompt, timeout=300, streaming=False):
|
||||||
|
"""
|
||||||
|
Non-streaming request (backward compatible).
|
||||||
|
Returns complete response string.
|
||||||
|
"""
|
||||||
|
if streaming:
|
||||||
|
raise ValueError("Use request_stream() for streaming requests")
|
||||||
return self.call(
|
return self.call(
|
||||||
system=system, prompt=prompt, timeout=timeout
|
system=system, prompt=prompt, streaming=False, timeout=timeout
|
||||||
).response
|
).response
|
||||||
|
|
||||||
|
def request_stream(self, system, prompt, timeout=300):
|
||||||
|
"""
|
||||||
|
Streaming request generator.
|
||||||
|
Yields response chunks as they arrive.
|
||||||
|
Usage:
|
||||||
|
for chunk in client.request_stream(system, prompt):
|
||||||
|
print(chunk.response, end='', flush=True)
|
||||||
|
"""
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
id = str(uuid.uuid4())
|
||||||
|
request = TextCompletionRequest(
|
||||||
|
system=system, prompt=prompt, streaming=True
|
||||||
|
)
|
||||||
|
|
||||||
|
end_time = time.time() + timeout
|
||||||
|
self.producer.send(request, properties={"id": id})
|
||||||
|
|
||||||
|
# Collect responses until end_of_stream
|
||||||
|
while time.time() < end_time:
|
||||||
|
try:
|
||||||
|
msg = self.consumer.receive(timeout_millis=2500)
|
||||||
|
except Exception:
|
||||||
|
continue
|
||||||
|
|
||||||
|
mid = msg.properties()["id"]
|
||||||
|
|
||||||
|
if mid == id:
|
||||||
|
value = msg.value()
|
||||||
|
|
||||||
|
# Handle errors
|
||||||
|
if value.error:
|
||||||
|
self.consumer.acknowledge(msg)
|
||||||
|
if value.error.type == "llm-error":
|
||||||
|
raise LlmError(value.error.message)
|
||||||
|
else:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"{value.error.type}: {value.error.message}"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.consumer.acknowledge(msg)
|
||||||
|
yield value
|
||||||
|
|
||||||
|
# Check if this is the final chunk
|
||||||
|
if getattr(value, 'end_of_stream', True):
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
# Ignore messages with wrong ID
|
||||||
|
self.consumer.acknowledge(msg)
|
||||||
|
|
||||||
|
if time.time() >= end_time:
|
||||||
|
raise TimeoutError("Timed out waiting for response")
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -19,7 +19,8 @@ class PromptRequestTranslator(MessageTranslator):
|
||||||
|
|
||||||
return PromptRequest(
|
return PromptRequest(
|
||||||
id=data.get("id"),
|
id=data.get("id"),
|
||||||
terms=terms
|
terms=terms,
|
||||||
|
streaming=data.get("streaming", False)
|
||||||
)
|
)
|
||||||
|
|
||||||
def from_pulsar(self, obj: PromptRequest) -> Dict[str, Any]:
|
def from_pulsar(self, obj: PromptRequest) -> Dict[str, Any]:
|
||||||
|
|
@ -51,4 +52,6 @@ class PromptResponseTranslator(MessageTranslator):
|
||||||
|
|
||||||
def from_response_with_completion(self, obj: PromptResponse) -> Tuple[Dict[str, Any], bool]:
|
def from_response_with_completion(self, obj: PromptResponse) -> Tuple[Dict[str, Any], bool]:
|
||||||
"""Returns (response_dict, is_final)"""
|
"""Returns (response_dict, is_final)"""
|
||||||
return self.from_pulsar(obj), True
|
# Check end_of_stream field to determine if this is the final message
|
||||||
|
is_final = getattr(obj, 'end_of_stream', True)
|
||||||
|
return self.from_pulsar(obj), is_final
|
||||||
|
|
@ -9,7 +9,8 @@ class TextCompletionRequestTranslator(MessageTranslator):
|
||||||
def to_pulsar(self, data: Dict[str, Any]) -> TextCompletionRequest:
|
def to_pulsar(self, data: Dict[str, Any]) -> TextCompletionRequest:
|
||||||
return TextCompletionRequest(
|
return TextCompletionRequest(
|
||||||
system=data["system"],
|
system=data["system"],
|
||||||
prompt=data["prompt"]
|
prompt=data["prompt"],
|
||||||
|
streaming=data.get("streaming", False)
|
||||||
)
|
)
|
||||||
|
|
||||||
def from_pulsar(self, obj: TextCompletionRequest) -> Dict[str, Any]:
|
def from_pulsar(self, obj: TextCompletionRequest) -> Dict[str, Any]:
|
||||||
|
|
@ -39,4 +40,6 @@ class TextCompletionResponseTranslator(MessageTranslator):
|
||||||
|
|
||||||
def from_response_with_completion(self, obj: TextCompletionResponse) -> Tuple[Dict[str, Any], bool]:
|
def from_response_with_completion(self, obj: TextCompletionResponse) -> Tuple[Dict[str, Any], bool]:
|
||||||
"""Returns (response_dict, is_final)"""
|
"""Returns (response_dict, is_final)"""
|
||||||
return self.from_pulsar(obj), True
|
# Check end_of_stream field to determine if this is the final message
|
||||||
|
is_final = getattr(obj, 'end_of_stream', True)
|
||||||
|
return self.from_pulsar(obj), is_final
|
||||||
|
|
@ -1,5 +1,5 @@
|
||||||
|
|
||||||
from pulsar.schema import Record, String, Array, Double, Integer
|
from pulsar.schema import Record, String, Array, Double, Integer, Boolean
|
||||||
|
|
||||||
from ..core.topic import topic
|
from ..core.topic import topic
|
||||||
from ..core.primitives import Error
|
from ..core.primitives import Error
|
||||||
|
|
@ -11,6 +11,7 @@ from ..core.primitives import Error
|
||||||
class TextCompletionRequest(Record):
|
class TextCompletionRequest(Record):
|
||||||
system = String()
|
system = String()
|
||||||
prompt = String()
|
prompt = String()
|
||||||
|
streaming = Boolean() # Default false for backward compatibility
|
||||||
|
|
||||||
class TextCompletionResponse(Record):
|
class TextCompletionResponse(Record):
|
||||||
error = Error()
|
error = Error()
|
||||||
|
|
@ -18,6 +19,7 @@ class TextCompletionResponse(Record):
|
||||||
in_token = Integer()
|
in_token = Integer()
|
||||||
out_token = Integer()
|
out_token = Integer()
|
||||||
model = String()
|
model = String()
|
||||||
|
end_of_stream = Boolean() # Indicates final message in stream
|
||||||
|
|
||||||
############################################################################
|
############################################################################
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
from pulsar.schema import Record, String, Map
|
from pulsar.schema import Record, String, Map, Boolean
|
||||||
|
|
||||||
from ..core.primitives import Error
|
from ..core.primitives import Error
|
||||||
from ..core.topic import topic
|
from ..core.topic import topic
|
||||||
|
|
@ -24,6 +24,9 @@ class PromptRequest(Record):
|
||||||
# JSON encoded values
|
# JSON encoded values
|
||||||
terms = Map(String())
|
terms = Map(String())
|
||||||
|
|
||||||
|
# Streaming support (default false for backward compatibility)
|
||||||
|
streaming = Boolean()
|
||||||
|
|
||||||
class PromptResponse(Record):
|
class PromptResponse(Record):
|
||||||
|
|
||||||
# Error case
|
# Error case
|
||||||
|
|
@ -35,4 +38,7 @@ class PromptResponse(Record):
|
||||||
# JSON encoded
|
# JSON encoded
|
||||||
object = String()
|
object = String()
|
||||||
|
|
||||||
|
# Indicates final message in stream
|
||||||
|
end_of_stream = Boolean()
|
||||||
|
|
||||||
############################################################################
|
############################################################################
|
||||||
|
|
@ -6,17 +6,63 @@ and user prompt. Both arguments are required.
|
||||||
import argparse
|
import argparse
|
||||||
import os
|
import os
|
||||||
import json
|
import json
|
||||||
from trustgraph.api import Api
|
import uuid
|
||||||
|
import asyncio
|
||||||
|
from websockets.asyncio.client import connect
|
||||||
|
|
||||||
default_url = os.getenv("TRUSTGRAPH_URL", 'http://localhost:8088/')
|
default_url = os.getenv("TRUSTGRAPH_URL", 'ws://localhost:8088/')
|
||||||
|
|
||||||
def query(url, flow_id, system, prompt):
|
async def query(url, flow_id, system, prompt, streaming=True):
|
||||||
|
|
||||||
api = Api(url).flow().id(flow_id)
|
if not url.endswith("/"):
|
||||||
|
url += "/"
|
||||||
|
|
||||||
resp = api.text_completion(system=system, prompt=prompt)
|
url = url + "api/v1/socket"
|
||||||
|
|
||||||
print(resp)
|
mid = str(uuid.uuid4())
|
||||||
|
|
||||||
|
async with connect(url) as ws:
|
||||||
|
|
||||||
|
req = {
|
||||||
|
"id": mid,
|
||||||
|
"service": "text-completion",
|
||||||
|
"flow": flow_id,
|
||||||
|
"request": {
|
||||||
|
"system": system,
|
||||||
|
"prompt": prompt,
|
||||||
|
"streaming": streaming
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
await ws.send(json.dumps(req))
|
||||||
|
|
||||||
|
while True:
|
||||||
|
|
||||||
|
msg = await ws.recv()
|
||||||
|
|
||||||
|
obj = json.loads(msg)
|
||||||
|
|
||||||
|
if "error" in obj:
|
||||||
|
raise RuntimeError(obj["error"])
|
||||||
|
|
||||||
|
if obj["id"] != mid:
|
||||||
|
continue
|
||||||
|
|
||||||
|
if "response" in obj["response"]:
|
||||||
|
if streaming:
|
||||||
|
# Stream output to stdout without newline
|
||||||
|
print(obj["response"]["response"], end="", flush=True)
|
||||||
|
else:
|
||||||
|
# Non-streaming: print complete response
|
||||||
|
print(obj["response"]["response"])
|
||||||
|
|
||||||
|
if obj["complete"]:
|
||||||
|
if streaming:
|
||||||
|
# Add final newline after streaming
|
||||||
|
print()
|
||||||
|
break
|
||||||
|
|
||||||
|
await ws.close()
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
|
|
||||||
|
|
@ -49,16 +95,23 @@ def main():
|
||||||
help=f'Flow ID (default: default)'
|
help=f'Flow ID (default: default)'
|
||||||
)
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'--no-streaming',
|
||||||
|
action='store_true',
|
||||||
|
help='Disable streaming (default: streaming enabled)'
|
||||||
|
)
|
||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|
||||||
query(
|
asyncio.run(query(
|
||||||
url=args.url,
|
url=args.url,
|
||||||
flow_id = args.flow_id,
|
flow_id=args.flow_id,
|
||||||
system=args.system[0],
|
system=args.system[0],
|
||||||
prompt=args.prompt[0],
|
prompt=args.prompt[0],
|
||||||
)
|
streaming=not args.no_streaming
|
||||||
|
))
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -10,20 +10,76 @@ using key=value arguments on the command line, and these replace
|
||||||
import argparse
|
import argparse
|
||||||
import os
|
import os
|
||||||
import json
|
import json
|
||||||
from trustgraph.api import Api
|
import uuid
|
||||||
|
import asyncio
|
||||||
|
from websockets.asyncio.client import connect
|
||||||
|
|
||||||
default_url = os.getenv("TRUSTGRAPH_URL", 'http://localhost:8088/')
|
default_url = os.getenv("TRUSTGRAPH_URL", 'ws://localhost:8088/')
|
||||||
|
|
||||||
def query(url, flow_id, template_id, variables):
|
async def query(url, flow_id, template_id, variables, streaming=True):
|
||||||
|
|
||||||
api = Api(url).flow().id(flow_id)
|
if not url.endswith("/"):
|
||||||
|
url += "/"
|
||||||
|
|
||||||
resp = api.prompt(id=template_id, variables=variables)
|
url = url + "api/v1/socket"
|
||||||
|
|
||||||
if isinstance(resp, str):
|
mid = str(uuid.uuid4())
|
||||||
print(resp)
|
|
||||||
else:
|
async with connect(url) as ws:
|
||||||
print(json.dumps(resp, indent=4))
|
|
||||||
|
req = {
|
||||||
|
"id": mid,
|
||||||
|
"service": "prompt",
|
||||||
|
"flow": flow_id,
|
||||||
|
"request": {
|
||||||
|
"id": template_id,
|
||||||
|
"variables": variables,
|
||||||
|
"streaming": streaming
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
await ws.send(json.dumps(req))
|
||||||
|
|
||||||
|
full_response = {"text": "", "object": ""}
|
||||||
|
|
||||||
|
while True:
|
||||||
|
|
||||||
|
msg = await ws.recv()
|
||||||
|
|
||||||
|
obj = json.loads(msg)
|
||||||
|
|
||||||
|
if "error" in obj:
|
||||||
|
raise RuntimeError(obj["error"])
|
||||||
|
|
||||||
|
if obj["id"] != mid:
|
||||||
|
continue
|
||||||
|
|
||||||
|
response = obj["response"]
|
||||||
|
|
||||||
|
# Handle text responses (streaming)
|
||||||
|
if "text" in response and response["text"]:
|
||||||
|
if streaming:
|
||||||
|
# Stream output to stdout without newline
|
||||||
|
print(response["text"], end="", flush=True)
|
||||||
|
full_response["text"] += response["text"]
|
||||||
|
else:
|
||||||
|
# Non-streaming: print complete response
|
||||||
|
print(response["text"])
|
||||||
|
|
||||||
|
# Handle object responses (JSON, never streamed)
|
||||||
|
if "object" in response and response["object"]:
|
||||||
|
full_response["object"] = response["object"]
|
||||||
|
|
||||||
|
if obj["complete"]:
|
||||||
|
if streaming and full_response["text"]:
|
||||||
|
# Add final newline after streaming text
|
||||||
|
print()
|
||||||
|
elif full_response["object"]:
|
||||||
|
# Print JSON object (pretty-printed)
|
||||||
|
print(json.dumps(json.loads(full_response["object"]), indent=4))
|
||||||
|
break
|
||||||
|
|
||||||
|
await ws.close()
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
|
|
||||||
|
|
@ -59,6 +115,12 @@ def main():
|
||||||
specified multiple times''',
|
specified multiple times''',
|
||||||
)
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'--no-streaming',
|
||||||
|
action='store_true',
|
||||||
|
help='Disable streaming (default: streaming enabled for text responses)'
|
||||||
|
)
|
||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
variables = {}
|
variables = {}
|
||||||
|
|
@ -73,12 +135,13 @@ specified multiple times''',
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|
||||||
query(
|
asyncio.run(query(
|
||||||
url=args.url,
|
url=args.url,
|
||||||
flow_id=args.flow_id,
|
flow_id=args.flow_id,
|
||||||
template_id=args.id[0],
|
template_id=args.id[0],
|
||||||
variables=variables,
|
variables=variables,
|
||||||
)
|
streaming=not args.no_streaming
|
||||||
|
))
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -32,7 +32,7 @@ from vertexai.generative_models import (
|
||||||
from anthropic import AnthropicVertex, RateLimitError
|
from anthropic import AnthropicVertex, RateLimitError
|
||||||
|
|
||||||
from .... exceptions import TooManyRequests
|
from .... exceptions import TooManyRequests
|
||||||
from .... base import LlmService, LlmResult
|
from .... base import LlmService, LlmResult, LlmChunk
|
||||||
|
|
||||||
# Module logger
|
# Module logger
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
@ -239,6 +239,123 @@ class Processor(LlmService):
|
||||||
logger.error(f"VertexAI LLM exception: {e}", exc_info=True)
|
logger.error(f"VertexAI LLM exception: {e}", exc_info=True)
|
||||||
raise e
|
raise e
|
||||||
|
|
||||||
|
def supports_streaming(self):
|
||||||
|
"""VertexAI supports streaming for both Gemini and Claude models"""
|
||||||
|
return True
|
||||||
|
|
||||||
|
async def generate_content_stream(self, system, prompt, model=None, temperature=None):
|
||||||
|
"""
|
||||||
|
Stream content generation from VertexAI (Gemini or Claude).
|
||||||
|
Yields LlmChunk objects with is_final=True on the last chunk.
|
||||||
|
"""
|
||||||
|
# Use provided model or fall back to default
|
||||||
|
model_name = model or self.default_model
|
||||||
|
# Use provided temperature or fall back to default
|
||||||
|
effective_temperature = temperature if temperature is not None else self.temperature
|
||||||
|
|
||||||
|
logger.debug(f"Using model (streaming): {model_name}")
|
||||||
|
logger.debug(f"Using temperature: {effective_temperature}")
|
||||||
|
|
||||||
|
try:
|
||||||
|
if 'claude' in model_name.lower():
|
||||||
|
# Claude/Anthropic streaming
|
||||||
|
logger.debug(f"Streaming request to Anthropic model '{model_name}'...")
|
||||||
|
client = self._get_anthropic_client()
|
||||||
|
|
||||||
|
total_in_tokens = 0
|
||||||
|
total_out_tokens = 0
|
||||||
|
|
||||||
|
with client.messages.stream(
|
||||||
|
model=model_name,
|
||||||
|
system=system,
|
||||||
|
messages=[{"role": "user", "content": prompt}],
|
||||||
|
max_tokens=self.api_params['max_output_tokens'],
|
||||||
|
temperature=effective_temperature,
|
||||||
|
top_p=self.api_params['top_p'],
|
||||||
|
top_k=self.api_params['top_k'],
|
||||||
|
) as stream:
|
||||||
|
# Stream text chunks
|
||||||
|
for text in stream.text_stream:
|
||||||
|
yield LlmChunk(
|
||||||
|
text=text,
|
||||||
|
in_token=None,
|
||||||
|
out_token=None,
|
||||||
|
model=model_name,
|
||||||
|
is_final=False
|
||||||
|
)
|
||||||
|
|
||||||
|
# Get final message with token counts
|
||||||
|
final_message = stream.get_final_message()
|
||||||
|
total_in_tokens = final_message.usage.input_tokens
|
||||||
|
total_out_tokens = final_message.usage.output_tokens
|
||||||
|
|
||||||
|
# Send final chunk with token counts
|
||||||
|
yield LlmChunk(
|
||||||
|
text="",
|
||||||
|
in_token=total_in_tokens,
|
||||||
|
out_token=total_out_tokens,
|
||||||
|
model=model_name,
|
||||||
|
is_final=True
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info(f"Input Tokens: {total_in_tokens}")
|
||||||
|
logger.info(f"Output Tokens: {total_out_tokens}")
|
||||||
|
|
||||||
|
else:
|
||||||
|
# Gemini streaming
|
||||||
|
logger.debug(f"Streaming request to Gemini model '{model_name}'...")
|
||||||
|
full_prompt = system + "\n\n" + prompt
|
||||||
|
|
||||||
|
llm, generation_config = self._get_gemini_model(model_name, effective_temperature)
|
||||||
|
|
||||||
|
response = llm.generate_content(
|
||||||
|
full_prompt,
|
||||||
|
generation_config=generation_config,
|
||||||
|
safety_settings=self.safety_settings,
|
||||||
|
stream=True # Enable streaming
|
||||||
|
)
|
||||||
|
|
||||||
|
total_in_tokens = 0
|
||||||
|
total_out_tokens = 0
|
||||||
|
|
||||||
|
# Stream chunks
|
||||||
|
for chunk in response:
|
||||||
|
if chunk.text:
|
||||||
|
yield LlmChunk(
|
||||||
|
text=chunk.text,
|
||||||
|
in_token=None,
|
||||||
|
out_token=None,
|
||||||
|
model=model_name,
|
||||||
|
is_final=False
|
||||||
|
)
|
||||||
|
|
||||||
|
# Accumulate token counts if available
|
||||||
|
if hasattr(chunk, 'usage_metadata') and chunk.usage_metadata:
|
||||||
|
if hasattr(chunk.usage_metadata, 'prompt_token_count'):
|
||||||
|
total_in_tokens = chunk.usage_metadata.prompt_token_count
|
||||||
|
if hasattr(chunk.usage_metadata, 'candidates_token_count'):
|
||||||
|
total_out_tokens = chunk.usage_metadata.candidates_token_count
|
||||||
|
|
||||||
|
# Send final chunk with token counts
|
||||||
|
yield LlmChunk(
|
||||||
|
text="",
|
||||||
|
in_token=total_in_tokens,
|
||||||
|
out_token=total_out_tokens,
|
||||||
|
model=model_name,
|
||||||
|
is_final=True
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info(f"Input Tokens: {total_in_tokens}")
|
||||||
|
logger.info(f"Output Tokens: {total_out_tokens}")
|
||||||
|
|
||||||
|
except (google.api_core.exceptions.ResourceExhausted, RateLimitError) as e:
|
||||||
|
logger.warning(f"Hit rate limit during streaming: {e}")
|
||||||
|
raise TooManyRequests()
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"VertexAI streaming exception: {e}", exc_info=True)
|
||||||
|
raise e
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def add_args(parser):
|
def add_args(parser):
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue