Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 9 additions & 0 deletions .env.example
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
129 changes: 129 additions & 0 deletions mingli_bench/models/astraflow_client.py
Original file line number Diff line number Diff line change
@@ -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),
)
17 changes: 16 additions & 1 deletion mingli_bench/models/factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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',
}

Expand All @@ -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
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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}")
Expand Down Expand Up @@ -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'
Expand Down
18 changes: 18 additions & 0 deletions mingli_bench/utils/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"),
Expand Down