From 6d175d4172322689b7537672051304e365769011 Mon Sep 17 00:00:00 2001 From: ucloudnb666 Date: Thu, 14 May 2026 17:12:02 +0800 Subject: [PATCH] feat: add Astraflow provider support Signed-off-by: ucloudnb666 --- .env.example | 9 ++ mingli_bench/models/astraflow_client.py | 129 ++++++++++++++++++++++++ mingli_bench/models/factory.py | 17 +++- mingli_bench/utils/config.py | 18 ++++ 4 files changed, 172 insertions(+), 1 deletion(-) create mode 100644 mingli_bench/models/astraflow_client.py diff --git a/.env.example b/.env.example index 58f3793..f5925bf 100644 --- a/.env.example +++ b/.env.example @@ -27,6 +27,15 @@ GOOGLE_API_KEY=your_google_api_key DEEPSEEK_API_KEY=your_deepseek_api_key DEEPSEEK_BASE_URL=https://api.deepseek.com +# Astraflow (by UCloud) — OpenAI-compatible platform supporting 200+ models +# Global endpoint — sign up at https://astraflow.ucloud-global.com +ASTRAFLOW_API_KEY=your_astraflow_api_key +# ASTRAFLOW_BASE_URL=https://api-us-ca.umodelverse.ai/v1 + +# Astraflow China endpoint — sign up at https://astraflow.ucloud.cn +# ASTRAFLOW_CN_API_KEY=your_astraflow_cn_api_key +# ASTRAFLOW_CN_BASE_URL=https://api.modelverse.cn/v1 + # Doubao / ByteDance DOUBAO_API_KEY=your_doubao_api_key DOUBAO_BASE_URL=your_doubao_base_url diff --git a/mingli_bench/models/astraflow_client.py b/mingli_bench/models/astraflow_client.py new file mode 100644 index 0000000..c6c01a2 --- /dev/null +++ b/mingli_bench/models/astraflow_client.py @@ -0,0 +1,129 @@ +""" +Astraflow model client implementation. +Astraflow (by UCloud) is an OpenAI-compatible platform supporting 200+ models. +Global endpoint : https://api-us-ca.umodelverse.ai/v1 (ASTRAFLOW_API_KEY) +China endpoint : https://api.modelverse.cn/v1 (ASTRAFLOW_CN_API_KEY) +""" + +import os +from typing import Optional +from openai import OpenAI + +from .base import ModelClient +from ..utils.logger import get_logger + +logger = get_logger(__name__) + + +class AstraflowClient(ModelClient): + """Client for Astraflow models (OpenAI-compatible API — global endpoint).""" + + # API key environment variables + API_KEY_ENV_VARS = ["ASTRAFLOW_API_KEY"] + + # Default base URL for Astraflow global endpoint + DEFAULT_BASE_URL = "https://api-us-ca.umodelverse.ai/v1" + + def __init__( + self, + model_name: str = "deepseek-ai/DeepSeek-R1", + api_key: Optional[str] = None, + base_url: Optional[str] = None, + **kwargs, + ): + """ + Initialize Astraflow client (global endpoint). + + Args: + model_name: Model to use (any model supported by Astraflow). + api_key: Astraflow API key (falls back to ASTRAFLOW_API_KEY env var). + base_url: API base URL (default: global endpoint). + **kwargs: Additional configuration. + """ + super().__init__( + model_name=model_name, + api_key=api_key, + api_key_env_vars=self.API_KEY_ENV_VARS, + **kwargs, + ) + + self.client = OpenAI( + api_key=self.api_key, + base_url=base_url or os.getenv("ASTRAFLOW_BASE_URL", self.DEFAULT_BASE_URL), + ) + + def generate(self, prompt: str, **kwargs) -> str: + """ + Generate response using the Astraflow API. + + Args: + prompt: Input prompt. + **kwargs: Override generation parameters. + + Returns: + Generated text. + """ + try: + params = self.get_generation_params(**kwargs) + + response = self.client.chat.completions.create( + model=self.model_name, + messages=[ + {"role": "system", "content": self.SYSTEM_PROMPT}, + {"role": "user", "content": prompt}, + ], + **params, + ) + + return response.choices[0].message.content.strip() + + except Exception as e: + self.handle_api_error("Astraflow generation", e) + raise + + def validate_api_key(self) -> bool: + """Validate Astraflow API key.""" + try: + self.client.models.list() + return True + except Exception as e: + logger.error(f"Invalid Astraflow API key: {e}") + return False + + +class AstraflowCNClient(AstraflowClient): + """Client for Astraflow models (OpenAI-compatible API — China endpoint).""" + + API_KEY_ENV_VARS = ["ASTRAFLOW_CN_API_KEY"] + DEFAULT_BASE_URL = "https://api.modelverse.cn/v1" + + def __init__( + self, + model_name: str = "deepseek-ai/DeepSeek-R1", + api_key: Optional[str] = None, + base_url: Optional[str] = None, + **kwargs, + ): + """ + Initialize Astraflow client (China endpoint). + + Args: + model_name: Model to use. + api_key: Astraflow CN API key (falls back to ASTRAFLOW_CN_API_KEY env var). + base_url: API base URL (default: China endpoint). + **kwargs: Additional configuration. + """ + # Call grandparent (ModelClient) directly so we can set our own env vars + # before the OpenAI client is constructed. + ModelClient.__init__( + self, + model_name=model_name, + api_key=api_key, + api_key_env_vars=self.API_KEY_ENV_VARS, + **kwargs, + ) + + self.client = OpenAI( + api_key=self.api_key, + base_url=base_url or os.getenv("ASTRAFLOW_CN_BASE_URL", self.DEFAULT_BASE_URL), + ) diff --git a/mingli_bench/models/factory.py b/mingli_bench/models/factory.py index ec55cf6..e0d7d0f 100644 --- a/mingli_bench/models/factory.py +++ b/mingli_bench/models/factory.py @@ -19,6 +19,8 @@ 'deepseek': 'pip install openai', 'anthropic': 'pip install anthropic', 'google': 'pip install google-generativeai', + 'astraflow': 'pip install openai', + 'astraflow_cn': 'pip install openai', 'doubao': 'pip install requests', } @@ -35,6 +37,8 @@ class ModelFactory: 'deepseek': ('.deepseek_client', 'DeepSeekClient'), 'doubao': ('.doubao_client', 'DoubaoClient'), 'openrouter': ('.openai_client', 'OpenAIClient'), # OpenAI-compatible API + 'astraflow': ('.astraflow_client', 'AstraflowClient'), # OpenAI-compatible API (global) + 'astraflow_cn': ('.astraflow_client', 'AstraflowCNClient'), # OpenAI-compatible API (China) } @classmethod @@ -105,6 +109,8 @@ def get_provider(cls, model_name: str) -> Optional[str]: return 'deepseek' if model_name.startswith('doubao-'): return 'doubao' + if model_name.startswith('astraflow-'): + return 'astraflow' return None @@ -144,7 +150,8 @@ def create(cls, raise ValueError( f"Cannot determine provider for model '{model_name}'. " f"Supported patterns: gpt-*, o1-*, o3-*, o4-*, claude-*, gemini-*, deepseek-*, doubao-*, " - f"or use OpenRouter format: provider/model-name (e.g., openai/gpt-4, nvidia/llama-3)" + f"or use OpenRouter format: provider/model-name (e.g., openai/gpt-4, nvidia/llama-3), " + f"or astraflow-* for Astraflow models" ) logger.info(f"Determined provider: {provider} for model: {model_name}") @@ -192,6 +199,14 @@ def list_supported_models(cls) -> Dict[str, list]: 'google': ['gemini-pro', 'gemini-1.5-pro', 'gemini-1.5-flash'], 'deepseek': ['deepseek-chat', 'deepseek-coder'], 'doubao': ['doubao-pro', 'doubao-lite'], + 'astraflow': [ + 'astraflow-deepseek-ai/DeepSeek-R1', 'astraflow-deepseek-ai/DeepSeek-V3', + 'astraflow-meta-llama/Llama-4-Maverick', + ], + 'astraflow_cn': [ + 'astraflow-deepseek-ai/DeepSeek-R1', 'astraflow-deepseek-ai/DeepSeek-V3', + 'astraflow-meta-llama/Llama-4-Maverick', + ], 'openrouter': [ 'openai/gpt-4', 'anthropic/claude-3-sonnet', 'google/gemini-2.0-flash', 'x-ai/grok-4', 'moonshotai/kimi-k2', 'deepseek/deepseek-r1' diff --git a/mingli_bench/utils/config.py b/mingli_bench/utils/config.py index 86a57eb..f4c9071 100644 --- a/mingli_bench/utils/config.py +++ b/mingli_bench/utils/config.py @@ -85,6 +85,24 @@ def load_config(env_file: Optional[str] = None) -> Dict[str, Any]: "max_tokens": default_max_tokens, }, + # Astraflow configuration (OpenAI-compatible, global endpoint) + # Sign up at https://astraflow.ucloud-global.com + "astraflow": { + "api_key": os.getenv("ASTRAFLOW_API_KEY"), + "base_url": os.getenv("ASTRAFLOW_BASE_URL", "https://api-us-ca.umodelverse.ai/v1"), + "temperature": default_temperature, + "max_tokens": default_max_tokens, + }, + + # Astraflow configuration (OpenAI-compatible, China endpoint) + # Sign up at https://astraflow.ucloud.cn + "astraflow_cn": { + "api_key": os.getenv("ASTRAFLOW_CN_API_KEY"), + "base_url": os.getenv("ASTRAFLOW_CN_BASE_URL", "https://api.modelverse.cn/v1"), + "temperature": default_temperature, + "max_tokens": default_max_tokens, + }, + # Doubao configuration "doubao": { "api_key": os.getenv("DOUBAO_API_KEY"),