Files
cmhub/apps/ai/providers/utils.py
T

264 lines
7.7 KiB
Python
Raw Normal View History

2026-07-02 10:33:15 +08:00
from __future__ import annotations
import base64
import binascii
import re
from pathlib import Path
from typing import Any, Iterable
from urllib.parse import urljoin, urlparse
import requests
from .base import AiProviderConfigError, AiResponseParseError
API_AUTO = "auto"
API_CHAT = "chat"
API_GEMINI = "gemini"
API_IMAGES = "images"
API_IMAGES_EDITS = "images_edits"
SUPPORTED_API_TYPES = {
API_AUTO,
API_CHAT,
API_GEMINI,
API_IMAGES,
API_IMAGES_EDITS,
}
RESOLUTION_TIMEOUTS = {"512": 180, "1K": 240, "2K": 360, "4K": 600}
BASE64_KEYS = {"image_base64", "base64", "b64_json", "data"}
IMAGE_EXTENSIONS = {".png", ".jpg", ".jpeg", ".webp", ".gif"}
TITLE_LEAD = re.compile(r"^\s*(?:\d+\s*[\.\)、::]|[-*•])\s*")
TITLE_CIRCLED = "①②③④⑤⑥⑦⑧⑨⑩"
TITLE_QUOTES = "\"'「」『』“”‘’"
TITLE_SPLIT = re.compile(r"[,,\r\n]+")
def resolution_timeout(resolution: str, default: int = 240) -> int:
return RESOLUTION_TIMEOUTS.get(str(resolution).strip().upper(), default)
def request_timeout(connect_timeout: int, read_timeout: int, resolution: str) -> tuple[int, int]:
resolved_read_timeout = read_timeout if read_timeout > 0 else resolution_timeout(resolution)
return connect_timeout, resolved_read_timeout
def detect_api_type(url: str, api_type: str = API_AUTO) -> str:
if api_type and api_type != API_AUTO:
if api_type not in SUPPORTED_API_TYPES:
raise AiProviderConfigError(f"unsupported api_type: {api_type}")
return api_type
lower_url = str(url).lower()
if "generatecontent" in lower_url or "gemini" in lower_url:
return API_GEMINI
if "/images/edits" in lower_url:
return API_IMAGES_EDITS
if "/images" in lower_url:
return API_IMAGES
return API_CHAT
def normalize_api_url(url: str, api_type: str) -> str:
raw = str(url).strip()
if not raw:
return raw
endpoint = {
API_CHAT: "chat/completions",
API_IMAGES: "images/generations",
API_IMAGES_EDITS: "images/edits",
}.get(api_type)
lower_path = urlparse(raw).path.lower().rstrip("/")
if api_type == API_GEMINI:
if "generatecontent" in lower_path:
return raw
return join_url(raw, "v1beta/models/{model}:generateContent")
if endpoint is None:
return raw
if lower_path.endswith("/" + endpoint):
return raw
if lower_path.endswith("/v1"):
return join_url(raw, endpoint)
return join_url(raw, "v1/" + endpoint)
def join_url(base_url: str, suffix: str) -> str:
base = str(base_url).rstrip("/") + "/"
return urljoin(base, suffix)
def image_bytes_to_data_url(image: bytes, mime_type: str = "image/png") -> str:
encoded = base64.b64encode(image).decode("ascii")
return f"data:{mime_type};base64,{encoded}"
def split_data_url(data_url: str) -> tuple[str, str]:
prefix, encoded = data_url.split(",", 1)
mime_type = prefix[len("data:") :].split(";", 1)[0]
return mime_type, encoded
def decode_image_data_url(data_url: str) -> bytes:
marker = ";base64,"
if marker not in data_url:
raise AiResponseParseError("unsupported data URL image format")
return base64.b64decode(data_url.split(marker, 1)[1])
def resolution_to_size(resolution: str) -> str:
2026-07-02 14:47:22 +08:00
key = str(resolution).strip().upper()
2026-07-02 10:33:15 +08:00
mapping = {
"512": "512x512",
2026-07-02 14:47:22 +08:00
"512PX": "512x512",
2026-07-02 10:33:15 +08:00
"1K": "1024x1024",
"2K": "2048x2048",
"4K": "4096x4096",
}
2026-07-02 14:47:22 +08:00
return mapping.get(key, str(resolution).strip())
2026-07-02 10:33:15 +08:00
def extract_image_from_response(
data: Any,
*,
session: requests.Session | None = None,
timeout: int = 60,
) -> bytes | None:
for key, value in walk_json_items(data):
if not isinstance(value, str):
continue
text = value.strip()
if text.startswith("data:image/"):
return decode_image_data_url(text)
if key and key.lower() in BASE64_KEYS and looks_like_base64(text):
try:
return base64.b64decode(text)
except (TypeError, ValueError, binascii.Error):
pass
for _key, value in walk_json_items(data):
if isinstance(value, str) and is_image_url(value):
return download_image_url(value, session=session, timeout=timeout)
return None
def download_image_url(
url: str,
*,
session: requests.Session | None = None,
timeout: int = 60,
) -> bytes:
client = session or requests.Session()
if hasattr(client, "trust_env"):
client.trust_env = False
response = client.get(url, timeout=timeout)
response.raise_for_status()
return response.content
def extract_titles_from_response(data: Any) -> tuple[str, ...]:
return tuple(clean_titles(extract_raw_text(data)))
def extract_text_from_response(data: Any) -> str:
titles = extract_titles_from_response(data)
return titles[0] if titles else ""
def extract_raw_text(data: Any) -> str:
if not isinstance(data, dict):
return ""
choices = data.get("choices")
if isinstance(choices, list) and choices and isinstance(choices[0], dict):
message = choices[0].get("message")
if isinstance(message, dict):
text = content_to_text(message.get("content"))
if text.strip():
return text
legacy = choices[0].get("text")
if isinstance(legacy, str) and legacy.strip():
return legacy
candidates = data.get("candidates")
if isinstance(candidates, list) and candidates and isinstance(candidates[0], dict):
content = candidates[0].get("content")
if isinstance(content, dict) and isinstance(content.get("parts"), list):
texts = [
part.get("text")
for part in content["parts"]
if isinstance(part, dict) and isinstance(part.get("text"), str)
]
joined = "\n".join(text for text in texts if text)
if joined.strip():
return joined
return ""
def content_to_text(content: Any) -> str:
if isinstance(content, str):
return content
if isinstance(content, list):
texts = []
for part in content:
if isinstance(part, dict) and isinstance(part.get("text"), str):
texts.append(part["text"])
elif isinstance(part, str):
texts.append(part)
return "\n".join(texts)
return ""
def clean_titles(text: str) -> list[str]:
if not text:
return []
out = []
for piece in TITLE_SPLIT.split(str(text)):
cleaned = clean_title_line(piece)
if cleaned:
out.append(cleaned)
return out
def clean_title_line(line: str) -> str:
stripped = line.strip()
if not stripped:
return ""
stripped = TITLE_LEAD.sub("", stripped)
stripped = stripped.lstrip(TITLE_CIRCLED).strip()
stripped = stripped.strip(TITLE_QUOTES).strip()
return stripped
def walk_json_items(value: Any, key: str | None = None) -> Iterable[tuple[str | None, Any]]:
yield key, value
if isinstance(value, dict):
for child_key, child_value in value.items():
yield from walk_json_items(child_value, str(child_key))
elif isinstance(value, list):
for child_value in value:
yield from walk_json_items(child_value, key)
def looks_like_base64(text: str) -> bool:
if len(text) < 8:
return False
try:
base64.b64decode(text, validate=True)
return True
except (TypeError, ValueError, binascii.Error):
return False
def is_image_url(text: str) -> bool:
parsed = urlparse(text.strip())
if parsed.scheme.lower() not in ("http", "https"):
return False
suffix = Path(parsed.path).suffix.lower()
return suffix in IMAGE_EXTENSIONS