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
42 changes: 42 additions & 0 deletions main.py
Original file line number Diff line number Diff line change
Expand Up @@ -682,6 +682,11 @@ def reload_env_globals():

CHAT_MODELS = model_list("CHAT_MODELS", CHAT_MODEL, ["gpt-4o-mini", "gemini-3.1-flash-image-preview-2k"])
IMAGE_MODELS = model_list("IMAGE_MODELS", IMAGE_MODEL, ["nano-banana-pro"])
ATLASCLOUD_DEFAULT_BASE_URL = "https://api.atlascloud.ai/v1"
ATLASCLOUD_DEFAULT_CHAT_MODELS = [
"qwen/qwen3.5-flash",
"deepseek-ai/deepseek-v4-pro",
]
VIDEO_MODELS = model_list("VIDEO_MODELS", "veo3-fast", [
# —— Veo 系列 ——
"veo2", "veo2-fast", "veo2-pro",
Expand All @@ -705,6 +710,8 @@ def reload_env_globals():
def provider_key_env(provider_id):
if provider_id == "comfly":
return "COMFLY_API_KEY"
if provider_id == "atlascloud":
return "ATLASCLOUD_API_KEY"
if provider_id == "modelscope":
return "MODELSCOPE_API_KEY"
if provider_id == "runninghub":
Expand Down Expand Up @@ -745,6 +752,8 @@ def provider_env_key_value(provider_id: str) -> str:
key = os.getenv(env_key, "") or read_api_env_value(env_key)
if key:
return key
if provider_id == "atlascloud":
return os.getenv("ATLAS_CLOUD_API_KEY", "") or read_api_env_value("ATLAS_CLOUD_API_KEY")
if provider_id == "modelscope":
return MODELSCOPE_API_KEY or ""
return ""
Expand Down Expand Up @@ -803,6 +812,22 @@ def default_api_providers():
"ms_loras": MODELSCOPE_DEFAULT_LORAS,
"ms_defaults_version": MODELSCOPE_DEFAULTS_VERSION,
},
{
"id": "atlascloud",
"name": "Atlas Cloud",
"base_url": ATLASCLOUD_DEFAULT_BASE_URL,
"protocol": "openai",
"image_request_mode": "openai",
"image_generation_endpoint": "",
"image_edit_endpoint": "",
"enabled": True,
"primary": False,
"image_models": [],
"chat_models": ATLASCLOUD_DEFAULT_CHAT_MODELS,
"video_models": [],
"ms_loras": [],
"ms_defaults_version": 0,
},
{
"id": "runninghub",
"name": "RunningHub",
Expand Down Expand Up @@ -862,6 +887,23 @@ def merge_default_api_providers(providers, inject_missing=True):
current["chat_models"] = chat_models
current["ms_loras"] = loras
current["ms_defaults_version"] = MODELSCOPE_DEFAULTS_VERSION
atlas_default = next((d for d in default_api_providers() if d["id"] == "atlascloud"), None)
if atlas_default:
current = next((item for item in merged if item.get("id") == "atlascloud"), None)
if not current:
if inject_missing:
merged.append(atlas_default)
else:
if not current.get("base_url"):
current["base_url"] = atlas_default["base_url"]
if not current.get("protocol"):
current["protocol"] = "openai"
current["image_models"] = model_list_from_values(current.get("image_models") or [])
current["chat_models"] = model_list_from_values([
*ATLASCLOUD_DEFAULT_CHAT_MODELS,
*(current.get("chat_models") or []),
])
current["video_models"] = model_list_from_values(current.get("video_models") or [])
rh_default = load_static_runninghub_provider() or next((d for d in default_api_providers() if d["id"] == "runninghub"), None)
if rh_default:
current = next((item for item in merged if item.get("id") == "runninghub"), None)
Expand Down
88 changes: 88 additions & 0 deletions tests/test_atlascloud_provider.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,88 @@
import json
import os
import sys
import tempfile
import unittest
from pathlib import Path
from unittest.mock import patch

sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
import main


class AtlasCloudProviderTests(unittest.TestCase):
def setUp(self):
self.temp = tempfile.TemporaryDirectory()
self.root = Path(self.temp.name)
self.providers_file = self.root / "api_providers.json"
self.api_env_file = self.root / ".env"
self.patches = [
patch.object(main, "API_PROVIDERS_FILE", str(self.providers_file)),
patch.object(main, "API_ENV_FILE", str(self.api_env_file)),
]
for item in self.patches:
item.start()

def tearDown(self):
for item in reversed(self.patches):
item.stop()
self.temp.cleanup()

def atlas_provider(self):
providers = main.load_api_providers()
return next(item for item in providers if item["id"] == "atlascloud")

def test_default_provider_is_registered_for_chat_models(self):
provider = self.atlas_provider()

self.assertEqual(provider["name"], "Atlas Cloud")
self.assertEqual(provider["base_url"], "https://api.atlascloud.ai/v1")
self.assertEqual(provider["protocol"], "openai")
self.assertEqual(provider["chat_models"][:2], ["qwen/qwen3.5-flash", "deepseek-ai/deepseek-v4-pro"])
self.assertEqual(provider["image_models"], [])
self.assertEqual(provider["video_models"], [])
self.assertEqual(main.public_provider(provider)["key_env"], "ATLASCLOUD_API_KEY")

def test_existing_atlas_provider_keeps_custom_models_with_safe_defaults(self):
self.providers_file.write_text(
json.dumps([
{
"id": "atlascloud",
"name": "Atlas",
"base_url": "",
"protocol": "openai",
"chat_models": ["custom/atlas-model"],
"image_models": [],
"video_models": [],
}
]),
encoding="utf-8",
)

provider = self.atlas_provider()

self.assertEqual(provider["base_url"], "https://api.atlascloud.ai/v1")
self.assertEqual(provider["chat_models"], [
"qwen/qwen3.5-flash",
"deepseek-ai/deepseek-v4-pro",
"custom/atlas-model",
])

def test_api_key_aliases_are_used_for_openai_compatible_chat(self):
with patch.dict(os.environ, {"ATLASCLOUD_API_KEY": "", "ATLAS_CLOUD_API_KEY": "alias-token"}, clear=False):
self.assertEqual(main.provider_env_key_value("atlascloud"), "alias-token")
base, headers, model = main.resolve_chat_provider("atlascloud", "", "")

self.assertEqual(base, "https://api.atlascloud.ai/v1")
self.assertEqual(model, "qwen/qwen3.5-flash")
self.assertEqual(headers["Authorization"], "Bearer alias-token")

def test_api_env_file_supports_atlas_cloud_alias(self):
self.api_env_file.write_text("ATLAS_CLOUD_API_KEY=file-alias\n", encoding="utf-8")

with patch.dict(os.environ, {"ATLASCLOUD_API_KEY": "", "ATLAS_CLOUD_API_KEY": ""}, clear=False):
self.assertEqual(main.provider_env_key_value("atlascloud"), "file-alias")


if __name__ == "__main__":
unittest.main()