| 1 | import os |
| 2 | import sys |
| 3 | |
| 4 | models_dir = os.path.dirname(os.path.abspath(__file__)) |
| 5 | backend_dir = os.path.dirname(models_dir) |
| 6 | if backend_dir not in sys.path: |
| 7 | sys.path.insert(0, backend_dir) |
| 8 | |
| 9 | import logging |
| 10 | import os |
| 11 | from typing import List, Optional |
| 12 | from config import Config |
| 13 | |
| 14 | try: |
| 15 | from models.vlm_dashscope import QwenVLClient |
| 16 | from models.vlm_gemini import GeminiVLClient |
| 17 | from models.vlm_gpt import GPTVLClient |
| 18 | except ImportError: |
| 19 | from vlm_dashscope import QwenVLClient |
| 20 | from vlm_gemini import GeminiVLClient |
| 21 | from vlm_gpt import GPTVLClient |
| 22 | |
| 23 | logger = logging.getLogger(__name__) |
| 24 | |
| 25 | |
| 26 | class VLM: |
| 27 | def __init__(self, |
| 28 | dashscope_api_key: Optional[str] = None, |
| 29 | dashscope_base_url: Optional[str] = None, |
| 30 | gemini_api_key: Optional[str] = None, |
| 31 | gemini_base_url: Optional[str] = None, |
| 32 | gpt_api_key: Optional[str] = None, |
| 33 | gpt_base_url: Optional[str] = None, |
| 34 | proxy: Optional[str] = None): |
| 35 | """ |
| 36 | Unified VLM (Vision Language Model) Client |
| 37 | Routes requests to DashScope (QwenVL) or Gemini based on model name. |
| 38 | """ |
| 39 | self._dashscope_api_key = dashscope_api_key |
| 40 | self._dashscope_base_url = dashscope_base_url |
| 41 | self._gemini_api_key = gemini_api_key |
| 42 | self._gemini_base_url = gemini_base_url |
| 43 | self._gpt_api_key = gpt_api_key |
| 44 | self._gpt_base_url = gpt_base_url |
| 45 | self._proxy = Config.provider_proxy("openai") if proxy is None else proxy |
| 46 | |
| 47 | self._dashscope_client = None |
| 48 | self._gemini_client = None |
| 49 | self._gpt_client = None |
| 50 | |
| 51 | @property |
| 52 | def dashscope_client(self): |
| 53 | if self._dashscope_client is None: |
| 54 | self._dashscope_client = QwenVLClient( |
| 55 | api_key=self._dashscope_api_key, |
| 56 | base_url=self._dashscope_base_url, |
| 57 | ) |
| 58 | return self._dashscope_client |
| 59 | |
| 60 | @property |
| 61 | def gemini_client(self): |
| 62 | if self._gemini_client is None: |
| 63 | self._gemini_client = GeminiVLClient( |
| 64 | api_key=self._gemini_api_key, |
| 65 | base_url=self._gemini_base_url, |
| 66 | ) |
| 67 | return self._gemini_client |
| 68 | |
| 69 | @property |
| 70 | def gpt_client(self): |
| 71 | if self._gpt_client is None: |
| 72 | self._gpt_client = GPTVLClient( |
| 73 | api_key=self._gpt_api_key, |
| 74 | base_url=self._gpt_base_url, |
| 75 | proxy=self._proxy, |
| 76 | ) |
| 77 | return self._gpt_client |
| 78 | |
| 79 | def query(self, |
| 80 | prompt: str, |
| 81 | image_paths: Optional[List[str]] = None, |
| 82 | model: str = "qwen3.6-plus", |
| 83 | session_id: Optional[str] = None) -> str: |
| 84 | if Config.PRINT_MODEL_INPUT: |
| 85 | lines = [ |
| 86 | "---- VLM REQUEST ----", |
| 87 | f"Prompt: {prompt}", |
| 88 | ] |
| 89 | if image_paths: |
| 90 | lines.append(f"Images: {len(image_paths)}") |
| 91 | for p in image_paths: |
| 92 | lines.append(" - [Base64图片]" if p.startswith("data:") else f" - {p}") |
| 93 | lines.append(f"Model: {model}") |
| 94 | if session_id: |
| 95 | lines.append(f"Session ID: {session_id}") |
| 96 | lines.append("-" * 30) |
| 97 | logger.info("\n%s", "\n".join(lines)) |
| 98 | |
| 99 | # Determine backend provider |
| 100 | model_lower = model.lower() |
| 101 | is_gemini = "gemini" in model_lower |
| 102 | is_gpt = "gpt" in model_lower |
| 103 | |
| 104 | if is_gemini: |
| 105 | # 处理图片路径 |
| 106 | processed_images = [] |
| 107 | for p in image_paths or []: |
| 108 | if p.startswith("data:") or p.startswith("http") or p.startswith("file://"): |
| 109 | processed_images.append(p) |
| 110 | else: |
| 111 | processed_images.append(p) # 传递原始路径,内部会处理 |
| 112 | return self.gemini_client.chat(text=prompt, images=processed_images, model=model) |
| 113 | elif is_gpt: |
| 114 | return self.gpt_client.chat(text=prompt, images=image_paths or [], model=model) |
| 115 | else: |
| 116 | # DashScope (Qwen/Kimi) - 需要将 base64 保存为临时文件 |
| 117 | file_urls = [] |
| 118 | import tempfile |
| 119 | import base64 as b64 |
| 120 | |
| 121 | for p in image_paths or []: |
| 122 | if p.startswith("data:"): |
| 123 | # Base64 数据 URL,需要解码并保存为临时文件 |
| 124 | try: |
| 125 | # 解析 data URL: data:image/png;base64,xxxxx |
| 126 | header, b64_data = p.split(",", 1) |
| 127 | mime_type = header.split(";")[0].replace("data:", "") |
| 128 | image_data = b64.b64decode(b64_data) |
| 129 | |
| 130 | # 创建临时文件 |
| 131 | suffix = f".{mime_type.split('/')[-1]}" if '/' in mime_type else ".png" |
| 132 | with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as tmp: |
| 133 | tmp.write(image_data) |
| 134 | temp_path = tmp.name |
| 135 | |
| 136 | abs_path = os.path.abspath(temp_path) |
| 137 | file_urls.append(f"file://{abs_path}") |
| 138 | except Exception as e: |
| 139 | logger.exception("Failed to process base64 image") |
| 140 | raise ValueError(f"无法解析 base64 图片: {e}") |
| 141 | elif p.startswith("http") or p.startswith("file://"): |
| 142 | file_urls.append(p) |
| 143 | else: |
| 144 | abs_path = os.path.abspath(p) |
| 145 | file_urls.append(f"file://{abs_path}") |
| 146 | return self.dashscope_client.chat(text=prompt, images=file_urls, model=model, stream=False) |
| 147 |