This commit is contained in:
9qeklajc
2025-11-24 22:24:43 +01:00
parent 95ffc612ca
commit 5eb4a40395
2 changed files with 23 additions and 32 deletions
+1 -2
View File
@@ -24,7 +24,7 @@ class BaseAPIClient(ABC):
pass pass
@abstractmethod @abstractmethod
async def generate_content_stream( def generate_content_stream(
self, self,
model: str, model: str,
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
@@ -32,7 +32,6 @@ class BaseAPIClient(ABC):
max_tokens: int | None = None, max_tokens: int | None = None,
**kwargs: Any, **kwargs: Any,
) -> AsyncGenerator[dict[str, Any], None]: ) -> AsyncGenerator[dict[str, Any], None]:
"""Generate content with streaming."""
pass pass
@abstractmethod @abstractmethod
+22 -30
View File
@@ -1,6 +1,6 @@
from __future__ import annotations from __future__ import annotations
from typing import Any, AsyncGenerator, Awaitable, Callable from typing import Any, AsyncGenerator
from openai import AsyncOpenAI from openai import AsyncOpenAI
@@ -26,19 +26,15 @@ class GeminiClient(BaseAPIClient):
max_tokens: int | None = None, max_tokens: int | None = None,
**kwargs: Any, **kwargs: Any,
) -> dict[str, Any]: ) -> dict[str, Any]:
"""Generate content using Gemini API via OpenAI SDK (non-streaming).""" from openai import NOT_GIVEN
args = {
"model": model,
"messages": messages,
}
if temperature is not None:
args["temperature"] = temperature
if max_tokens is not None:
args["max_tokens"] = max_tokens
if "top_p" in kwargs:
args["top_p"] = kwargs["top_p"]
response = await self.client.chat.completions.create(**args) response = await self.client.chat.completions.create(
model=model,
messages=messages, # type: ignore
temperature=temperature if temperature is not None else NOT_GIVEN,
max_tokens=max_tokens if max_tokens is not None else NOT_GIVEN,
top_p=kwargs.get("top_p", NOT_GIVEN),
)
return response.model_dump() return response.model_dump()
async def generate_content_stream( async def generate_content_stream(
@@ -47,26 +43,22 @@ class GeminiClient(BaseAPIClient):
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
temperature: float | None = None, temperature: float | None = None,
max_tokens: int | None = None, max_tokens: int | None = None,
usage_callback: Callable[[dict[str, Any]], None] | None = None,
completion_callback: Callable[[str, dict[str, Any] | None], Awaitable[None]]
| None = None,
**kwargs: Any, **kwargs: Any,
) -> AsyncGenerator[dict[str, Any], None]: ) -> AsyncGenerator[dict[str, Any], None]:
"""Generate content using Gemini API via OpenAI SDK (streaming).""" from openai import NOT_GIVEN
args = {
"model": model,
"messages": messages,
"stream": True,
"stream_options": {"include_usage": True},
}
if temperature is not None:
args["temperature"] = temperature
if max_tokens is not None:
args["max_tokens"] = max_tokens
if "top_p" in kwargs:
args["top_p"] = kwargs["top_p"]
stream = await self.client.chat.completions.create(**args) usage_callback = kwargs.get("usage_callback")
completion_callback = kwargs.get("completion_callback")
stream = await self.client.chat.completions.create(
model=model,
messages=messages, # type: ignore
stream=True,
stream_options={"include_usage": True},
temperature=temperature if temperature is not None else NOT_GIVEN,
max_tokens=max_tokens if max_tokens is not None else NOT_GIVEN,
top_p=kwargs.get("top_p", NOT_GIVEN),
)
final_usage = None final_usage = None