Merge branch 'refs/heads/feature/gemini-llm' into update/enhance-gemini-with-proxy

This commit is contained in:
caixypromise
2025-05-12 04:38:09 +08:00
@@ -1,140 +1,133 @@
import google.generativeai as genai # core/providers/llm/gemini_sdk.py
from core.utils.util import check_model_key import os, json, uuid
from core.providers.llm.base import LLMProviderBase from types import SimpleNamespace
from config.logger import setup_logging from typing import Any, Dict, List
import requests
import json
from google import generativeai as genai
from google.generativeai import types, GenerationConfig
from core.providers.llm.base import LLMProviderBase
from core.utils.util import check_model_key
from config.logger import setup_logging
from google.generativeai.types import GenerateContentResponse
log = setup_logging()
TAG = __name__ TAG = __name__
logger = setup_logging()
class LLMProvider(LLMProviderBase): class LLMProvider(LLMProviderBase):
def __init__(self, config): def __init__(self, cfg: Dict[str, Any]):
"""初始化Gemini LLM Provider""" self.model_name = cfg.get("model_name", "gemini-2.0-flash")
self.model_name = config.get("model_name", "gemini-1.5-pro") self.api_key = cfg["api_key"]
self.api_key = config.get("api_key") proxy = cfg.get("https_proxy") or cfg.get("http_proxy")
self.http_proxy = config.get("http_proxy")
self.https_proxy = config.get("https_proxy")
have_key = check_model_key("LLM", self.api_key)
if not have_key: if not check_model_key("LLM", self.api_key):
return raise ValueError("无效的Gemini API Key,请检查是否配置正确")
try: if proxy:
# 初始化Gemini客户端 os.environ["HTTPS_PROXY"] = os.environ["HTTP_PROXY"] = proxy
# 配置代理(如果提供了代理配置) log.bind(tag=TAG).info(f"Gemini 代理地址: {proxy}")
self.proxies = None
if self.http_proxy is not "" or self.https_proxy is not "":
self.proxies = { genai.configure(api_key=self.api_key)
"http": self.http_proxy, self.model = genai.GenerativeModel(self.model_name)
"https": self.https_proxy,
}
logger.bind(tag=TAG).info(f"Gemini set proxys:{self.proxies}")
# 使用猴子补丁修改 google-generativeai 库的请求会话
# 使用 session 对象配置 genai self.gen_cfg = GenerationConfig(
temperature=0.7,
top_p=0.9,
top_k=40,
max_output_tokens=2048,
)
genai.configure(api_key=self.api_key)
self.model = genai.GenerativeModel(self.model_name)
# 设置生成参数 @staticmethod
self.generation_config = { def _build_tools(funcs: List[Dict[str, Any]] | None):
"temperature": 0.7, if not funcs:
"top_p": 0.9, return None
"top_k": 40, return [types.Tool(function_declarations=[
"max_output_tokens": 2048, types.FunctionDeclaration(
} name=f["function"]["name"],
self.chat = None description=f["function"]["description"],
except Exception as e: parameters=f["function"]["parameters"],
logger.bind(tag=TAG).error(f"Gemini初始化失败: {e}") )
self.model = None for f in funcs
])]
# Gemini文档提到,无需维护session-id,直接用dialogue拼接而成
def response(self, session_id, dialogue): def response(self, session_id, dialogue):
"""生成Gemini对话响应""" yield from self._generate(dialogue, None)
if not self.model:
yield "【Gemini服务未正确初始化】"
return
try:
# 处理对话历史
chat_history = []
for msg in dialogue[:-1]: # 历史对话
role = "model" if msg["role"] == "assistant" else "user"
content = msg["content"].strip()
if content:
chat_history.append({"role": role, "parts": [{"text": content}]})
# 获取当前消息
current_msg = dialogue[-1]["content"]
# 构建请求体
request_body = {
"contents": chat_history
+ [{"role": "user", "parts": [{"text": current_msg}]}],
"generationConfig": self.generation_config,
}
# 构建请求URL
url = f"https://generativelanguage.googleapis.com/v1beta/models/{self.model_name}:generateContent?key={self.api_key}"
# 构建请求头
headers = {
"Content-Type": "application/json",
}
# 发送POST请求,经测试手动 request 无法使用 stream 模式
if self.proxies:
response = requests.post(
url,
headers=headers,
json=request_body,
stream=False,
proxies=self.proxies,
)
try:
data = response.json() # 直接解析JSON
if "candidates" in data and data["candidates"]:
yield data["candidates"][0]["content"]["parts"][0]["text"]
else:
yield "未找到候选回复。"
except json.JSONDecodeError as e:
yield f"JSON解码错误:{e}"
except Exception as e:
yield f"发生错误:{e}"
else:
logger.bind(tag=TAG).info(f"Gemini stream mode ")
chat = self.model.start_chat(history=chat_history)
# 发送消息并获取流式响应
response = chat.send_message(
current_msg, stream=True, generation_config=self.generation_config
)
# 处理流式响应
for chunk in response:
if hasattr(chunk, "text") and chunk.text:
yield chunk.text
except Exception as e:
error_msg = str(e)
logger.bind(tag=TAG).error(f"Gemini响应生成错误: {error_msg}")
# 针对不同错误返回友好提示
if "Rate limit" in error_msg:
yield "【Gemini服务请求太频繁,请稍后再试】"
elif "Invalid API key" in error_msg:
yield "【Gemini API key无效】"
else:
yield f"【Gemini服务响应异常: {error_msg}"
except requests.exceptions.RequestException as e:
yield f"请求失败:{e}"
except json.JSONDecodeError as e:
yield f"JSON解码错误:{e}"
except Exception as e:
yield f"发生错误:{e}"
def response_with_functions(self, session_id, dialogue, functions=None): def response_with_functions(self, session_id, dialogue, functions=None):
logger.bind(tag=TAG).info(f"gemini暂未实现完整的工具调用(function call") yield from self._generate(dialogue, self._build_tools(functions))
return self.response(session_id, dialogue)
def _generate(self, dialogue, tools):
role_map = {"assistant": "model", "user": "user"}
contents: list = []
# 拼接对话
for m in dialogue:
r = m["role"]
if r == "assistant" and "tool_calls" in m:
tc = m["tool_calls"][0]
contents.append({
"role": "model",
"parts": [{"function_call": {
"name": tc["function"]["name"],
"args": json.loads(tc["function"]["arguments"]),
}}],
})
continue
if r == "tool":
contents.append({
"role": "model",
"parts": [{"text": str(m.get("content", ""))}],
})
continue
contents.append({
"role": role_map.get(r, "user"),
"parts": [{"text": str(m.get("content", ""))}],
})
stream: GenerateContentResponse = self.model.generate_content(
contents=contents,
generation_config=self.gen_cfg,
tools=tools,
stream=True,
)
try:
for chunk in stream:
cand = chunk.candidates[0]
for part in cand.content.parts:
# a) 函数调用-通常是最后一段话才是函数调用
if getattr(part, "function_call", None):
fc = part.function_call
yield None, [SimpleNamespace(
id=uuid.uuid4().hex,
type="function",
function=SimpleNamespace(
name=fc.name,
arguments=json.dumps(dict(fc.args),
ensure_ascii=False),
),
)]
return
# b) 普通文本
if getattr(part, "text", None):
yield part.text if tools is None else (part.text, None)
finally:
if tools is not None:
yield None, None # functionmode 结束,返回哑包
# 关闭stream,预留后续打断对话功能的功能方法,官方文档推荐打断对话要关闭上一个流,可以有效减少配额计费和资源占用
@staticmethod
def _safe_finish_stream(stream: GenerateContentResponse):
if hasattr(stream, "resolve"):
stream.resolve() # Gemini SDK version ≥ 0.5.0
elif hasattr(stream, "close"):
stream.close() # Gemini SDK version < 0.5.0
else:
for _ in stream: # 兜底耗尽
pass