diff --git a/main.py b/main.py index 6e36dba56..455866595 100644 --- a/main.py +++ b/main.py @@ -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", @@ -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": @@ -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 "" @@ -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", @@ -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) diff --git a/tests/test_atlascloud_provider.py b/tests/test_atlascloud_provider.py new file mode 100644 index 000000000..2ba5d99fa --- /dev/null +++ b/tests/test_atlascloud_provider.py @@ -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()