"""Application-level configuration for cmshopee. This module owns `data/config.json` plus `data/config/ai_models.json`. `config.json` stores local app settings and AI role/generation parameters. AI provider definitions and local plaintext API keys live in ignored `data/config/ai_models.json`. """ import copy import json import os import shutil import sys import tempfile import urllib.error import urllib.parse import urllib.request DATA_DIR_NAME = "data" def app_base_dir() -> str: """Return the application directory used as the parent of local data.""" if getattr(sys, "frozen", False): return os.path.dirname(os.path.abspath(sys.executable)) return os.path.dirname(os.path.dirname(os.path.abspath(__file__))) def default_data_dir(base_dir=None) -> str: """Return the default local user data directory.""" return os.path.abspath(os.path.join(base_dir or app_base_dir(), DATA_DIR_NAME)) CONFIG_PATH = os.path.join(default_data_dir(), "config.json") AI_MODELS_PATH = os.path.join(default_data_dir(), "config", "ai_models.json") CMHUB_CONFIG_PATH = os.path.join(default_data_dir(), "config", "cmhub.json") CATEGORIES = {"text", "image"} API_TYPES = {"chat", "images_edits", "auto"} AI_BACKENDS = {"direct", "cmhub"} AI_CONCURRENCY_MIN = 1 AI_CONCURRENCY_MAX = 5 AI_RETRY_MIN = 0 AI_RETRY_MAX = 10 RUNTIME_CONFIG_KEYS = { "config_path", "ai_models_path", "cmhub_config_path", "data_dir", } LEGACY_USER_DATA_PATHS = ( "config.json", "config", "cmshopee.db", "cmshopee.db-wal", "cmshopee.db-shm", "db.sqlite", "chrome_user_data_dir", "images", "logs", "prompts", "title_prompt.txt", ) DEFAULT_CONFIG = { "chrome_path": r"C:\Program Files\Google\Chrome\Application\chrome.exe", "user_data_root": "chrome_user_data_dir", "image_dir": "images", "db_path": "cmshopee.db", "default_debug_port": 9222, "debug_port_range": [9222, 9260], "cdp_ready_timeout": 60, "ai": { "default_text_model": "GPT-5.5 文本", "default_image_model": "Nano Banana 2", "generate_cover": False, "backend": "cmhub", "cmhub": { "base_url": "", "title_alias": "", "image_alias": "", "connect_timeout": 10, "use_system_proxy": False, "download_with_curl": "auto", "check_balance_before_batch": False, }, "title_concurrency": 4, "image_concurrency": 4, "retry": 2, "jpg_quality": 90, "resolution": "1k", "resolution_timeouts": { "512": 180, "1k": 240, "2k": 360, "4k": 600, }, }, "shopee_update": { "test_item_id": "51100639510", "allow_real_submit": False, "allow_cover_update": False, "max_items_per_run": 1, "close_success_tab": False, "dry_run": False, "parallel_accounts": False, "max_parallel_accounts": 2, }, } DEFAULT_AI_MODELS_CONFIG = { "models": [ { "name": "GPT-5.5 文本", "category": "text", "enabled": True, "url": "", "model": "", "api_key": "", "api_type": "chat", "connect_timeout_seconds": 30, "timeout_seconds": 0, "extra_body": {}, }, { "name": "Nano Banana 2", "category": "image", "enabled": True, "url": "https://api.vectorengine.ai/v1/chat/completions", "model": "gemini-3.1-flash-image-preview", "api_key": "", "api_type": "auto", "connect_timeout_seconds": 30, "timeout_seconds": 0, "extra_body": {}, }, ] } DEFAULT_CMHUB_CONFIG = { "api_key": "", } SECRET_FIELD_NAMES = {"api_key", "apikey", "key", "token", "password"} class ConfigError(RuntimeError): """Raised when app configuration is missing or malformed.""" class DataMigrationConflictError(ConfigError): """Raised when old-layout data cannot be moved into data/ safely.""" class DataDirectoryWriteError(ConfigError): """Raised when the local data directory is not writable.""" def _data_dir_for_config_path(path) -> str: config_path = os.path.abspath(path or CONFIG_PATH) if config_path == os.path.abspath(CONFIG_PATH): return default_data_dir() return os.path.dirname(config_path) def data_dir(config=None) -> str: """Return the absolute local user data directory.""" if isinstance(config, dict): if config.get("data_dir"): return os.path.abspath(str(config["data_dir"])) if config.get("config_path"): return _data_dir_for_config_path(config["config_path"]) return default_data_dir() def _strip_runtime_config_keys(config): if not isinstance(config, dict): return config return { key: copy.deepcopy(value) for key, value in config.items() if key not in RUNTIME_CONFIG_KEYS } def _runtime_paths(config_path, data_dir_path=None): root = os.path.abspath(data_dir_path or _data_dir_for_config_path(config_path)) return { "config_path": os.path.abspath(config_path or os.path.join(root, "config.json")), "data_dir": root, "ai_models_path": os.path.join(root, "config", "ai_models.json"), "cmhub_config_path": os.path.join(root, "config", "cmhub.json"), } def _with_runtime_paths(config, path): result = copy.deepcopy(config) result.update(_runtime_paths(path)) return result def _data_parent_dir(config=None) -> str: return os.path.dirname(data_dir(config)) def resolve_data_path(path, config=None) -> str: """Resolve a user-data path under data_dir unless it is already absolute.""" text = str(path or "").strip() if not text: return "" if os.path.isabs(text): return os.path.abspath(text) normalized = os.path.normpath(text) first_part = normalized.split(os.sep, 1)[0] if first_part == DATA_DIR_NAME: return os.path.abspath(os.path.join(_data_parent_dir(config), normalized)) return os.path.abspath(os.path.join(data_dir(config), normalized)) def data_path(*parts, config=None) -> str: return os.path.abspath(os.path.join(data_dir(config), *[str(part) for part in parts])) def ai_models_config_path(config=None) -> str: if isinstance(config, dict) and config.get("ai_models_path"): return os.path.abspath(str(config["ai_models_path"])) return data_path("config", "ai_models.json", config=config) def cmhub_config_file_path(config=None) -> str: if isinstance(config, dict) and config.get("cmhub_config_path"): return os.path.abspath(str(config["cmhub_config_path"])) return data_path("config", "cmhub.json", config=config) def title_prompt_path(config=None) -> str: return data_path("title_prompt.txt", config=config) def cover_prompts_dir(config=None) -> str: return data_path("prompts", "cover", config=config) def diagnostic_log_dir(config=None) -> str: return data_path("logs", config=config) def _is_within(path, parent) -> bool: try: return os.path.commonpath([os.path.abspath(path), os.path.abspath(parent)]) == os.path.abspath(parent) except ValueError: return False def migrate_legacy_user_data(base_dir=None, data_dir_path=None): """Move old exe-top-level local data into data_dir without overwriting.""" base = os.path.abspath(base_dir or app_base_dir()) target_root = os.path.abspath(data_dir_path or default_data_dir(base)) if base == target_root: return [] os.makedirs(target_root, exist_ok=True) moves = [] seen_sources = set() for relative in LEGACY_USER_DATA_PATHS: source = os.path.abspath(os.path.join(base, relative)) if source in seen_sources or not os.path.exists(source): continue seen_sources.add(source) if _is_within(source, target_root): continue if not _is_within(source, base): raise ConfigError(f"旧数据路径不在程序目录内: {source}") target = os.path.abspath(os.path.join(target_root, relative)) if not _is_within(target, target_root): raise ConfigError(f"迁移目标路径不在 data 目录内: {target}") moves.append((relative, source, target)) conflicts = [relative for relative, _source, target in moves if os.path.exists(target)] if conflicts: raise DataMigrationConflictError( "检测到旧布局数据和 data 目录内数据同时存在,无法自动迁移。" "请先手动合并或备份后再启动。冲突项: " + "、".join(conflicts) ) moved = [] for relative, source, target in moves: os.makedirs(os.path.dirname(target), exist_ok=True) shutil.move(source, target) moved.append(relative) return moved def ensure_writable_data_dir(path=None) -> str: root = os.path.abspath(path or default_data_dir()) try: os.makedirs(root, exist_ok=True) with tempfile.NamedTemporaryFile( mode="w", encoding="utf-8", prefix=".cmshopee_write_test_", dir=root, delete=False, ) as fh: marker = fh.name fh.write("ok") os.remove(marker) except Exception as exc: raise DataDirectoryWriteError( f"数据目录不可写:{root}。请把程序放到可写目录,勿放 Program Files。" ) from exc return root def prepare_data_dir(base_dir=None, data_dir_path=None, migrate=True) -> str: root = os.path.abspath(data_dir_path or default_data_dir(base_dir)) if migrate: migrate_legacy_user_data(base_dir=base_dir, data_dir_path=root) return ensure_writable_data_dir(root) def default_config() -> dict: """Return a new copy of the default config.""" return copy.deepcopy(DEFAULT_CONFIG) def default_ai_models_config() -> dict: """Return a new copy of the default AI model list.""" return copy.deepcopy(DEFAULT_AI_MODELS_CONFIG) def _deep_merge(defaults, loaded): if not isinstance(defaults, dict): return copy.deepcopy(loaded) if loaded is not None else copy.deepcopy(defaults) merged = copy.deepcopy(defaults) if not isinstance(loaded, dict): return merged for key, value in loaded.items(): if isinstance(merged.get(key), dict) and isinstance(value, dict): merged[key] = _deep_merge(merged[key], value) else: merged[key] = copy.deepcopy(value) return merged def _normalize_config_values(config): if not isinstance(config, dict): return config ai = config.get("ai") if isinstance(ai, dict): ai["title_concurrency"] = _clamp_int( ai.get("title_concurrency"), AI_CONCURRENCY_MIN, AI_CONCURRENCY_MAX, DEFAULT_CONFIG["ai"]["title_concurrency"], ) ai["image_concurrency"] = _clamp_int( ai.get("image_concurrency"), AI_CONCURRENCY_MIN, AI_CONCURRENCY_MAX, DEFAULT_CONFIG["ai"]["image_concurrency"], ) ai["retry"] = _clamp_int( ai.get("retry"), AI_RETRY_MIN, AI_RETRY_MAX, DEFAULT_CONFIG["ai"]["retry"], ) cmhub = ai.get("cmhub") if isinstance(cmhub, dict): cmhub["base_url"] = normalize_cmhub_base_url(cmhub.get("base_url", "")) cmhub["download_with_curl"] = _normalize_cmhub_download_with_curl( cmhub.get("download_with_curl", "auto") ) return config def _clamp_int(value, minimum, maximum, default): try: number = int(value) except (TypeError, ValueError): number = int(default) return min(int(maximum), max(int(minimum), number)) def _normalize_cmhub_download_with_curl(value): if isinstance(value, bool): return "true" if value else "false" text = str(value or "auto").strip().lower() if text in {"auto", "true", "false"}: return text return "auto" def _assert_no_secrets(config): def visit(value, path): if isinstance(value, dict): for key, child in value.items(): lowered = str(key).lower() if _is_secret_field_name(lowered): raise ConfigError( f"config.json 不允许保存敏感字段: {'.'.join(path + [str(key)])}" ) visit(child, path + [str(key)]) elif isinstance(value, list): for index, child in enumerate(value): visit(child, path + [str(index)]) visit(config, []) def _is_secret_field_name(name) -> bool: lowered = str(name).lower() return lowered in SECRET_FIELD_NAMES or lowered.endswith( ("_key", "_token", "_password") ) def mask_secret(secret) -> str: """Return a display-safe representation of a secret value.""" if secret is None: return "" if isinstance(secret, (dict, list, tuple, set)): return "***" if secret else "" text = str(secret or "") if not text: return "" if len(text) <= 8: return "***" return f"{text[:4]}***{text[-4:]}" def mask_email(email) -> str: """Return a display-safe email string with the local part masked.""" text = str(email or "").strip() if not text or "@" not in text: return text local, domain = text.split("@", 1) if not local or not domain: return mask_secret(text) if len(local) == 1: masked_local = "***" elif len(local) == 2: masked_local = f"{local[0]}***" else: masked_local = f"{local[0]}***{local[-1]}" return f"{masked_local}@{domain}" def sanitize_for_log(value): """Return a copy of value with secret fields masked for logs/status/export.""" if isinstance(value, dict): sanitized = {} for key, child in value.items(): lowered = str(key).lower() if _is_secret_field_name(key): sanitized[key] = mask_secret(child) elif lowered == "email" or lowered.endswith("_email"): sanitized[key] = mask_email(child) else: sanitized[key] = sanitize_for_log(child) return sanitized if isinstance(value, list): return [sanitize_for_log(item) for item in value] if isinstance(value, tuple): return tuple(sanitize_for_log(item) for item in value) return value def redact_secrets(text, secret_values=None) -> str: """Replace known secret values inside free-form text.""" redacted = str(text) for secret in secret_values or []: secret_text = str(secret or "") if secret_text: redacted = redacted.replace(secret_text, "***") return redacted def default_cmhub_config() -> dict: """Return a new copy of the default cmhub key config.""" return copy.deepcopy(DEFAULT_CMHUB_CONFIG) def _normalize_cmhub_config(config): if config is None: config = {} if not isinstance(config, dict): raise ConfigError("cmhub 配置必须是对象") return {"api_key": str(config.get("api_key", "") or "")} def load_cmhub_config(path=CMHUB_CONFIG_PATH) -> dict: """Load cmhub API key config. Missing file means key is not configured.""" if not os.path.exists(path): return default_cmhub_config() with open(path, "r", encoding="utf-8") as fh: try: loaded = json.load(fh) except json.JSONDecodeError as exc: raise ConfigError(f"cmhub 配置不是有效 JSON: {path}") from exc return _normalize_cmhub_config(loaded) def save_cmhub_config(config, path=CMHUB_CONFIG_PATH) -> dict: """Persist cmhub API key config, including the local plaintext key.""" normalized = _normalize_cmhub_config(config) directory = os.path.dirname(os.path.abspath(path)) if directory: os.makedirs(directory, exist_ok=True) with open(path, "w", encoding="utf-8") as fh: json.dump(normalized, fh, ensure_ascii=False, indent=2) fh.write("\n") return normalized def get_cmhub_api_key(path=CMHUB_CONFIG_PATH, masked=False) -> str: key = load_cmhub_config(path).get("api_key", "") return mask_secret(key) if masked else key def save_config(config, path=CONFIG_PATH) -> dict: """Persist config to JSON and return the normalized config.""" normalized = _deep_merge(DEFAULT_CONFIG, _strip_runtime_config_keys(config)) _normalize_config_values(normalized) _assert_no_secrets(normalized) directory = os.path.dirname(os.path.abspath(path)) if directory: os.makedirs(directory, exist_ok=True) with open(path, "w", encoding="utf-8") as fh: json.dump(normalized, fh, ensure_ascii=False, indent=2) fh.write("\n") return _with_runtime_paths(normalized, path) def load_config(path=CONFIG_PATH) -> dict: """Load config, writing defaults first if the file does not exist.""" if not os.path.exists(path): return save_config(default_config(), path=path) with open(path, "r", encoding="utf-8") as fh: try: loaded = json.load(fh) except json.JSONDecodeError as exc: raise ConfigError(f"配置文件不是有效 JSON: {path}") from exc normalized = _deep_merge(DEFAULT_CONFIG, loaded) _normalize_config_values(normalized) _assert_no_secrets(normalized) return _with_runtime_paths(normalized, path) def update_config(updates, path=CONFIG_PATH) -> dict: """Merge updates into the persisted config.""" config = load_config(path) return save_config(_deep_merge(config, updates), path=path) def _config_or_load(config): return load_config() if config is None else config def chrome_path(config=None) -> str: return _config_or_load(config).get("chrome_path", "") def user_data_root(config=None) -> str: cfg = _config_or_load(config) return resolve_data_path(cfg.get("user_data_root", "chrome_user_data_dir"), cfg) def image_dir(config=None) -> str: cfg = _config_or_load(config) return resolve_data_path(cfg.get("image_dir", "images"), cfg) def db_path(config=None) -> str: cfg = _config_or_load(config) return resolve_data_path(cfg.get("db_path", "cmshopee.db"), cfg) def default_debug_port(config=None) -> int: return int(_config_or_load(config).get("default_debug_port", 9222)) def debug_port_range(config=None) -> tuple: values = _config_or_load(config).get("debug_port_range", [9222, 9260]) if not isinstance(values, list) or len(values) != 2: raise ConfigError("debug_port_range 必须是 [start, end]") return int(values[0]), int(values[1]) def cdp_ready_timeout(config=None) -> int: return int(_config_or_load(config).get("cdp_ready_timeout", 60)) def ai_config(config=None) -> dict: return copy.deepcopy(_config_or_load(config).get("ai", DEFAULT_CONFIG["ai"])) def ai_backend(config=None) -> str: ai = ai_config(config) default_backend = DEFAULT_CONFIG["ai"]["backend"] backend = str(ai.get("backend", default_backend) or default_backend).strip().lower() if backend not in AI_BACKENDS: raise ConfigError("AI backend 必须是 direct 或 cmhub") return backend def cmhub_config(config=None) -> dict: ai = ai_config(config) value = ai.get("cmhub", {}) if not isinstance(value, dict): raise ConfigError("ai.cmhub 必须是对象") merged = _deep_merge(DEFAULT_CONFIG["ai"]["cmhub"], value) merged["base_url"] = normalize_cmhub_base_url(merged.get("base_url", "")) merged["title_alias"] = str(merged.get("title_alias", "") or "").strip() merged["image_alias"] = str(merged.get("image_alias", "") or "").strip() merged["connect_timeout"] = int(merged.get("connect_timeout", 10) or 10) merged["download_with_curl"] = _normalize_cmhub_download_with_curl( merged.get("download_with_curl", "auto") ) merged["check_balance_before_batch"] = bool(merged.get("check_balance_before_batch", False)) if merged["connect_timeout"] <= 0: raise ConfigError("ai.cmhub.connect_timeout 必须大于 0") return merged def normalize_cmhub_base_url(base_url) -> str: """Return the cmhub gateway root URL without path/query/fragment.""" text = str(base_url or "").strip() if not text: return "" parts = urllib.parse.urlsplit(text) if parts.scheme and parts.netloc: return urllib.parse.urlunsplit((parts.scheme, parts.netloc, "", "", "")) if parts.netloc: return urllib.parse.urlunsplit((parts.scheme, parts.netloc, "", "", "")) without_query = text.split("?", 1)[0].split("#", 1)[0].strip().strip("/") if "://" not in without_query: return without_query.split("/", 1)[0] return text.rstrip("/") def cmhub_request_url(base_url, endpoint) -> str: base = normalize_cmhub_base_url(base_url) path = "/" + str(endpoint or "").strip().lstrip("/") if not base: return path return base + path def response_timeout(config=None) -> int: ai = ai_config(config) resolution = str(ai.get("resolution", DEFAULT_CONFIG["ai"]["resolution"])) timeouts = ai.get("resolution_timeouts", {}) if resolution not in timeouts: raise ConfigError(f"未配置分辨率 {resolution} 的返回超时") return int(timeouts[resolution]) def _normalize_ai_model(model): if not isinstance(model, dict): raise ConfigError("AI 模型定义必须是对象") normalized = { "name": str(model.get("name", "")).strip(), "category": str(model.get("category", "")).strip(), "enabled": bool(model.get("enabled", True)), "url": str(model.get("url", "")).strip(), "model": str(model.get("model", "")).strip(), "api_key": str(model.get("api_key", "")), "api_type": str(model.get("api_type", "auto")).strip() or "auto", "connect_timeout_seconds": int(model.get("connect_timeout_seconds", 30) or 30), "timeout_seconds": int(model.get("timeout_seconds", 0) or 0), "extra_body": copy.deepcopy(model.get("extra_body", {})), } if not normalized["name"]: raise ConfigError("AI 模型 name 不能为空") if normalized["category"] not in CATEGORIES: raise ConfigError("AI 模型 category 必须是 text 或 image") if normalized["api_type"] not in API_TYPES: raise ConfigError("AI 模型 api_type 必须是 chat、images_edits 或 auto") if normalized["connect_timeout_seconds"] <= 0: raise ConfigError("AI 模型 connect_timeout_seconds 必须大于 0") if normalized["timeout_seconds"] < 0: raise ConfigError("AI 模型 timeout_seconds 不能小于 0") if not isinstance(normalized["extra_body"], dict): raise ConfigError("AI 模型 extra_body 必须是对象") return normalized def _normalize_ai_models_config(config): models = config.get("models") if isinstance(config, dict) else None if not isinstance(models, list): raise ConfigError("AI 模型清单必须包含 models 列表") normalized = {"models": [_normalize_ai_model(model) for model in models]} _assert_unique_model_names(normalized["models"]) _assert_required_categories(normalized["models"]) return normalized def _assert_unique_model_names(models): names = [model["name"] for model in models] duplicates = sorted({name for name in names if names.count(name) > 1}) if duplicates: raise ConfigError(f"AI 模型 name 重复: {', '.join(duplicates)}") def _assert_required_categories(models): enabled_categories = { model["category"] for model in models if model.get("enabled", True) } missing = sorted(CATEGORIES - enabled_categories) if missing: raise ConfigError( "AI 模型清单至少需要启用一个 text 和一个 image 模型,缺少: " + ", ".join(missing) ) def _model_index(models, name): for index, model in enumerate(models): if model["name"] == name: return index raise ConfigError(f"AI 模型不存在: {name}") def _mask_api_key(api_key): return mask_secret(api_key) def _public_model(model): public = copy.deepcopy(model) public["api_key"] = _mask_api_key(public.get("api_key", "")) public["api_key_set"] = bool(model.get("api_key")) return public def save_ai_models_config(config, path=AI_MODELS_PATH) -> dict: """Persist AI model definitions, including local plaintext API keys.""" normalized = _normalize_ai_models_config(config) directory = os.path.dirname(os.path.abspath(path)) if directory: os.makedirs(directory, exist_ok=True) with open(path, "w", encoding="utf-8") as fh: json.dump(normalized, fh, ensure_ascii=False, indent=2) fh.write("\n") return normalized def load_ai_models_config(path=AI_MODELS_PATH) -> dict: """Load AI model definitions, writing defaults first if missing.""" if not os.path.exists(path): return save_ai_models_config(default_ai_models_config(), path=path) with open(path, "r", encoding="utf-8") as fh: try: loaded = json.load(fh) except json.JSONDecodeError as exc: raise ConfigError(f"AI 模型清单不是有效 JSON: {path}") from exc return _normalize_ai_models_config(loaded) def list_ai_models(category=None, path=AI_MODELS_PATH, reveal_api_key=False): """Return AI models, optionally filtered by category.""" if category is not None and category not in CATEGORIES: raise ConfigError("category 必须是 text 或 image") models = load_ai_models_config(path)["models"] filtered = [ copy.deepcopy(model) for model in models if category is None or model["category"] == category ] if reveal_api_key: return filtered return [_public_model(model) for model in filtered] def add_ai_model(model, path=AI_MODELS_PATH) -> None: config = load_ai_models_config(path) normalized = _normalize_ai_model(model) if any(item["name"] == normalized["name"] for item in config["models"]): raise ConfigError(f"AI 模型 name 已存在: {normalized['name']}") config["models"].append(normalized) save_ai_models_config(config, path=path) def update_ai_model(model_name, path=AI_MODELS_PATH, **fields) -> None: if not fields: return config = load_ai_models_config(path) index = _model_index(config["models"], model_name) updated = copy.deepcopy(config["models"][index]) updated.update(fields) normalized = _normalize_ai_model(updated) if normalized["name"] != model_name and any( model["name"] == normalized["name"] for model in config["models"] ): raise ConfigError(f"AI 模型 name 已存在: {normalized['name']}") config["models"][index] = normalized save_ai_models_config(config, path=path) def delete_ai_model(name, path=AI_MODELS_PATH) -> None: config = load_ai_models_config(path) index = _model_index(config["models"], name) remaining = config["models"][:index] + config["models"][index + 1 :] _assert_required_categories(remaining) save_ai_models_config({"models": remaining}, path=path) def get_model(name, path=AI_MODELS_PATH) -> dict: """Return the model definition including api_key. Callers must not log it.""" models = load_ai_models_config(path)["models"] return copy.deepcopy(models[_model_index(models, name)]) def model_request_url(model) -> str: """Return the HTTP endpoint used for a configured AI model.""" raw_url = str(model.get("url", "") or "").strip() api_type = str(model.get("api_type", "auto") or "auto").strip() if api_type == "images_edits": return _append_default_endpoint(raw_url, "images/edits") return _append_default_endpoint(raw_url, "chat/completions") def _append_default_endpoint(raw_url, endpoint): if not raw_url: return raw_url parts = urllib.parse.urlsplit(raw_url) path = parts.path.rstrip("/") lowered = path.lower() endpoint_path = "/" + endpoint.strip("/") if lowered.endswith(endpoint_path): return raw_url base_markers = ("", "/v1", "/v1beta", "/api/v1", "/api/v1beta") if lowered in base_markers or lowered.endswith(base_markers[1:]): path = path + endpoint_path return urllib.parse.urlunsplit( (parts.scheme, parts.netloc, path, parts.query, parts.fragment) ) return raw_url def _test_request_payload(model): if model["api_type"] == "images_edits": payload = {"model": model["model"], "prompt": "ping"} payload.update(model.get("extra_body", {})) return payload payload = { "model": model["model"], "messages": [{"role": "user", "content": "ping"}], } payload.update(model.get("extra_body", {})) return payload def test_ai_model(name, path=AI_MODELS_PATH) -> dict: """Send a minimal request to the configured model endpoint.""" model = None try: model = get_model(name, path=path) missing = [ field for field in ("url", "model", "api_key") if not str(model.get(field, "")).strip() ] if missing: return {"ok": False, "error": "模型缺少字段: " + ", ".join(missing)} data = json.dumps(_test_request_payload(model), ensure_ascii=False).encode( "utf-8" ) request = urllib.request.Request( model_request_url(model), data=data, headers={ "Authorization": "Bearer " + model["api_key"], "Content-Type": "application/json", }, method="POST", ) with urllib.request.urlopen( request, timeout=int(model["connect_timeout_seconds"]) ) as response: response.read(1024) return {"ok": 200 <= response.status < 300, "status": response.status} except urllib.error.HTTPError as exc: return {"ok": False, "status": exc.code, "error": f"HTTP {exc.code}"} except urllib.error.URLError as exc: secret_values = [model.get("api_key")] if model else [] return {"ok": False, "error": redact_secrets(str(exc.reason), secret_values)} except Exception as exc: secret_values = [model.get("api_key")] if model else [] return {"ok": False, "error": redact_secrets(str(exc), secret_values)}