refactor: use shared http clients

This commit is contained in:
arabcoders 2026-01-26 17:42:24 +03:00
parent 41bca681e5
commit 98f25027bb
9 changed files with 527 additions and 370 deletions

View file

@ -1,12 +1,12 @@
# flake8: noqa: ARG004 # flake8: noqa: ARG004
from typing import Any from typing import TYPE_CHECKING, Any
import httpx
from yt_dlp.utils.networking import random_user_agent
from app.features.tasks.definitions.results import HandleTask, TaskFailure, TaskResult from app.features.tasks.definitions.results import HandleTask, TaskFailure, TaskResult
from app.library.config import Config from app.library.config import Config
from app.library.httpx_client import async_client from app.library.httpx_client import Globals, build_request_headers, get_async_client, resolve_curl_transport
if TYPE_CHECKING:
import httpx
class BaseHandler: class BaseHandler:
@ -33,7 +33,7 @@ class BaseHandler:
@staticmethod @staticmethod
async def request( async def request(
url: str, headers: dict | None = None, ytdlp_opts: dict | None = None, **kwargs url: str, headers: dict | None = None, ytdlp_opts: dict | None = None, **kwargs
) -> httpx.Response: ) -> "httpx.Response":
""" """
Make an HTTP request. Make an HTTP request.
@ -50,31 +50,21 @@ class BaseHandler:
headers = {} if not isinstance(headers, dict) else headers headers = {} if not isinstance(headers, dict) else headers
ytdlp_opts = {} if not isinstance(ytdlp_opts, dict) else ytdlp_opts ytdlp_opts = {} if not isinstance(ytdlp_opts, dict) else ytdlp_opts
opts: dict[str, Any] = { use_curl = resolve_curl_transport()
"headers": { request_headers = build_request_headers(
"User-Agent": random_user_agent(), base_headers=headers,
}, user_agent=Globals.get_random_agent(),
} use_curl=use_curl,
)
try: proxy = ytdlp_opts.get("proxy", None)
from httpx_curl_cffi import AsyncCurlTransport, CurlOpt client = get_async_client(proxy=proxy, use_curl=use_curl)
method = kwargs.pop("method", "GET").upper()
opts["transport"] = AsyncCurlTransport( timeout = ytdlp_opts.get("timeout", ytdlp_opts.get("socket_timeout", 120))
impersonate="chrome", return await client.request(
default_headers=True, method=method,
curl_options={CurlOpt.FRESH_CONNECT: True}, url=url,
) headers=request_headers,
opts["headers"].pop("User-Agent", None) timeout=timeout,
except Exception: **kwargs,
pass )
for k, v in headers.items():
opts["headers"][k] = v
if proxy := ytdlp_opts.get("proxy", None):
opts["proxy"] = proxy
async with async_client(**opts) as client:
method = kwargs.pop("method", "GET").upper()
timeout = ytdlp_opts.get("timeout", ytdlp_opts.get("socket_timeout", 120))
return await client.request(method=method, url=url, timeout=timeout, **kwargs)

View file

@ -14,7 +14,6 @@ from urllib.parse import urljoin
import jmespath import jmespath
from parsel import Selector from parsel import Selector
from yt_dlp.utils.networking import random_user_agent
from app.features.tasks.definitions.results import HandleTask, TaskFailure, TaskItem, TaskResult from app.features.tasks.definitions.results import HandleTask, TaskFailure, TaskItem, TaskResult
from app.features.tasks.definitions.schemas import ( from app.features.tasks.definitions.schemas import (
@ -23,7 +22,7 @@ from app.features.tasks.definitions.schemas import (
) )
from app.library.cache import Cache from app.library.cache import Cache
from app.library.config import Config from app.library.config import Config
from app.library.httpx_client import async_client from app.library.httpx_client import Globals, build_request_headers, get_async_client, resolve_curl_transport
from app.library.Utils import fetch_info, get_archive_id from app.library.Utils import fetch_info, get_archive_id
from ._base_handler import BaseHandler from ._base_handler import BaseHandler
@ -48,7 +47,7 @@ class GenericTaskHandler(BaseHandler):
"""Modification times of source files to detect changes.""" """Modification times of source files to detect changes."""
@classmethod @classmethod
async def refresh_definitions(cls, force: bool = False) -> None: async def refresh_definitions(cls, force: bool = False) -> list[TaskDefinition]:
""" """
Refresh the cached task definitions if source files have changed. Refresh the cached task definitions if source files have changed.
@ -66,12 +65,7 @@ class GenericTaskHandler(BaseHandler):
repo = TaskDefinitionsRepository.get_instance() repo = TaskDefinitionsRepository.get_instance()
models = await repo.list() models = await repo.list()
definitions: list[TaskDefinition] = [] cls._definitions = [model_to_schema(model) for model in models]
for model in models:
td = model_to_schema(model)
definitions.append(td)
cls._definitions = definitions
return cls._definitions return cls._definitions
except Exception as exc: except Exception as exc:
LOG.error(f"Failed to load task definitions from database: {exc}") LOG.error(f"Failed to load task definitions from database: {exc}")
@ -97,10 +91,10 @@ class GenericTaskHandler(BaseHandler):
try: try:
for matcher in definition.match_url: for matcher in definition.match_url:
pattern_str = None pattern_str: str | None = None
if matcher.startswith("/") and matcher.endswith("/") and len(matcher) > 2: if matcher.startswith("/") and matcher.endswith("/") and len(matcher) > 2:
pattern_str: str = matcher[1:-1] pattern_str = matcher[1:-1]
else: else:
pattern_str = fnmatch.translate(matcher) pattern_str = fnmatch.translate(matcher)
@ -274,53 +268,37 @@ class GenericTaskHandler(BaseHandler):
""" """
headers: dict[str, str] = {**definition.definition.request.headers} headers: dict[str, str] = {**definition.definition.request.headers}
client_options: dict[str, Any] = { use_curl = resolve_curl_transport()
"headers": { request_headers = build_request_headers(
"User-Agent": random_user_agent(), base_headers=headers,
} user_agent=Globals.get_random_agent(),
} use_curl=use_curl,
)
try:
from httpx_curl_cffi import AsyncCurlTransport, CurlOpt
client_options["transport"] = AsyncCurlTransport(
impersonate="chrome",
default_headers=True,
curl_options={CurlOpt.FRESH_CONNECT: True},
)
client_options["headers"].pop("User-Agent", None)
except Exception:
pass
if headers:
client_options["headers"].update(headers)
if proxy := ytdlp_opts.get("proxy"):
client_options["proxy"] = proxy
timeout_value: float | Any = definition.definition.request.timeout or ytdlp_opts.get("socket_timeout", 120) timeout_value: float | Any = definition.definition.request.timeout or ytdlp_opts.get("socket_timeout", 120)
async with async_client(**client_options) as client: client = get_async_client(proxy=ytdlp_opts.get("proxy"), use_curl=use_curl)
response: httpx.Response = await client.request( response: httpx.Response = await client.request(
method=definition.definition.request.method.upper(), method=definition.definition.request.method.upper(),
url=url, url=url,
params=definition.definition.request.params or None, params=definition.definition.request.params or None,
data=definition.definition.request.data, data=definition.definition.request.data,
json=definition.definition.request.json_data, json=definition.definition.request.json_data,
timeout=timeout_value, timeout=timeout_value,
) headers=request_headers,
response.raise_for_status() )
response.raise_for_status()
if "json" == definition.definition.response.type: if "json" == definition.definition.response.type:
try: try:
json_data: dict[str, Any] = response.json() json_data: dict[str, Any] = response.json()
except Exception as exc: except Exception as exc:
LOG.error(f"Failed to decode JSON response from '{url}': {exc}") LOG.error(f"Failed to decode JSON response from '{url}': {exc}")
return response.text, None return response.text, None
return response.text, json_data return response.text, json_data
return response.text, None return response.text, None
@staticmethod @staticmethod
async def _fetch_with_selenium( async def _fetch_with_selenium(

View file

@ -327,7 +327,8 @@ async def task_metadata(request: Request, repo: TasksRepository, config: Config,
The response object. The response object.
""" """
task_id = request.match_info.get("id") if not (task_id := request.match_info.get("id")):
return web.json_response(data={"error": "No task id."}, status=web.HTTPBadRequest.status_code)
try: try:
if not (model := await repo.get(int(task_id))): if not (model := await repo.get(int(task_id))):
@ -335,7 +336,6 @@ async def task_metadata(request: Request, repo: TasksRepository, config: Config,
data={"error": f"Task '{task_id}' does not exist."}, status=web.HTTPNotFound.status_code data={"error": f"Task '{task_id}' does not exist."}, status=web.HTTPNotFound.status_code
) )
# Convert to extended Task with handler methods
task = ExtendedTask.model_validate(model) task = ExtendedTask.model_validate(model)
(save_path, _) = get_file(config.download_path, task.folder) (save_path, _) = get_file(config.download_path, task.folder)
@ -348,6 +348,11 @@ async def task_metadata(request: Request, repo: TasksRepository, config: Config,
metadata, status, message = await task.fetch_metadata() metadata, status, message = await task.fetch_metadata()
if not status: if not status:
return web.json_response(data={"error": message}, status=web.HTTPBadRequest.status_code) return web.json_response(data={"error": message}, status=web.HTTPBadRequest.status_code)
if not isinstance(metadata, dict):
return web.json_response(
data={"error": "Failed to get metadata."},
status=web.HTTPBadRequest.status_code,
)
if not task.folder: if not task.folder:
try: try:
@ -395,7 +400,7 @@ async def task_metadata(request: Request, repo: TasksRepository, config: Config,
from app.yt_dlp_plugins.postprocessor.nfo_maker import NFOMakerPP from app.yt_dlp_plugins.postprocessor.nfo_maker import NFOMakerPP
title: str = sanitize_filename(info.get("title")) title: str = sanitize_filename(str(info.get("title") or ""))
info_file: Path = save_path / f"{title} [{info.get('id')}].info.json" info_file: Path = save_path / f"{title} [{info.get('id')}].info.json"
info_file.write_text(encoder.encode(metadata), encoding="utf-8") info_file.write_text(encoder.encode(metadata), encoding="utf-8")
info["json_file"] = str(info_file.relative_to(config.download_path)) info["json_file"] = str(info_file.relative_to(config.download_path))
@ -406,9 +411,7 @@ async def task_metadata(request: Request, repo: TasksRepository, config: Config,
xml_content = "<tvshow>\n" xml_content = "<tvshow>\n"
xml_content += f" <title>{NFOMakerPP._escape_text(info.get('title'))}</title>\n" xml_content += f" <title>{NFOMakerPP._escape_text(info.get('title'))}</title>\n"
if info.get("description"): if info.get("description"):
xml_content += ( xml_content += f" <plot>{NFOMakerPP._escape_text(NFOMakerPP._clean_description(str(info.get('description') or '')))}</plot>\n"
f" <plot>{NFOMakerPP._escape_text(NFOMakerPP._clean_description(info.get('description')))}</plot>\n"
)
if info.get("id"): if info.get("id"):
xml_content += f" <id>{NFOMakerPP._escape_text(info.get('id'))}</id>\n" xml_content += f" <id>{NFOMakerPP._escape_text(info.get('id'))}</id>\n"
if info.get("id_type") and info.get("id"): if info.get("id_type") and info.get("id"):
@ -425,52 +428,53 @@ async def task_metadata(request: Request, repo: TasksRepository, config: Config,
xml_file.write_text(xml_content, encoding="utf-8") xml_file.write_text(xml_content, encoding="utf-8")
try: try:
from yt_dlp.utils.networking import random_user_agent from app.library.httpx_client import (
Globals,
from app.library.httpx_client import async_client build_request_headers,
get_async_client,
resolve_curl_transport,
)
ytdlp_args: dict = task.get_ytdlp_opts().get_all() ytdlp_args: dict = task.get_ytdlp_opts().get_all()
opts: dict[str, Any] = { use_curl = resolve_curl_transport()
"headers": { request_headers = build_request_headers(
"User-Agent": request.headers.get("User-Agent", ytdlp_args.get("user_agent", random_user_agent())), user_agent=request.headers.get("User-Agent", ytdlp_args.get("user_agent", Globals.get_random_agent())),
}, use_curl=use_curl,
} )
if proxy := ytdlp_args.get("proxy"):
opts["proxy"] = proxy
try: client = get_async_client(proxy=ytdlp_args.get("proxy"), use_curl=use_curl)
from httpx_curl_cffi import AsyncCurlTransport thumbnails = info.get("thumbnails", {})
if not isinstance(thumbnails, dict):
opts["transport"] = AsyncCurlTransport( thumbnails = {}
impersonate="chrome", info["thumbnails"] = thumbnails
default_headers=True, for key in thumbnails:
) url: str | None = None
opts.pop("headers", None) try:
except Exception: url = thumbnails.get(key)
pass LOG.info(f"Fetching thumbnail '{key}' from '{url}'")
if not url:
async with async_client(**opts) as client:
for key in info.get("thumbnails", {}):
try:
url = info["thumbnails"][key]
LOG.info(f"Fetching thumbnail '{key}' from '{url}'")
if not url:
continue
try:
validate_url(url, allow_internal=config.allow_internal_urls)
except ValueError:
LOG.warning(f"Invalid thumbnail url '{url}'")
continue
resp = await client.request(method="GET", url=url, follow_redirects=True)
img_file = save_path / f"{key}.jpg"
img_file.write_bytes(resp.content)
info["thumbnails"][key] = str(img_file.relative_to(config.download_path))
except Exception as e:
LOG.warning(f"Failed to fetch thumbnail '{key}' from '{url}'. '{e!s}'")
continue continue
try:
validate_url(url, allow_internal=config.allow_internal_urls)
except ValueError:
LOG.warning(f"Invalid thumbnail url '{url}'")
continue
resp = await client.request(
method="GET",
url=url,
follow_redirects=True,
headers=request_headers,
)
img_file = save_path / f"{key}.jpg"
img_file.write_bytes(resp.content)
thumbnails[key] = str(img_file.relative_to(config.download_path))
except Exception as e:
url_log = url or "unknown"
LOG.warning(f"Failed to fetch thumbnail '{key}' from '{url_log}'. '{e!s}'")
continue
except Exception as e: except Exception as e:
LOG.warning(f"Failed to fetch thumbnails. '{e!s}'") LOG.warning(f"Failed to fetch thumbnails. '{e!s}'")

View file

@ -8,7 +8,7 @@ from aiohttp import web
from .cache import Cache from .cache import Cache
from .config import Config from .config import Config
from .Events import EventBus, Events from .Events import EventBus, Events
from .httpx_client import async_client from .httpx_client import get_async_client
from .Scheduler import Scheduler from .Scheduler import Scheduler
from .Singleton import Singleton from .Singleton import Singleton
from .version import APP_VERSION from .version import APP_VERSION
@ -147,36 +147,37 @@ class UpdateChecker(metaclass=Singleton):
try: try:
LOG.info(f"Checking for {name} updates...") LOG.info(f"Checking for {name} updates...")
async with async_client(timeout=10.0) as client: client = get_async_client(use_curl=False)
response = await client.get( response = await client.get(
api_url, api_url,
headers={"Accept": "application/vnd.github+json"}, headers={"Accept": "application/vnd.github+json"},
) timeout=10.0,
)
if 200 != response.status_code: if 200 != response.status_code:
LOG.warning(f"Failed to check for {name} updates: HTTP {response.status_code}") LOG.warning(f"Failed to check for {name} updates: HTTP {response.status_code}")
return ("error", None) return ("error", None)
data: dict[str, Any] = response.json() data: dict[str, Any] = response.json()
latest_tag: str = data.get("tag_name", "") latest_tag: str = data.get("tag_name", "")
if not latest_tag: if not latest_tag:
LOG.warning(f"No tag_name found in {name} GitHub release data.") LOG.warning(f"No tag_name found in {name} GitHub release data.")
return ("error", None) return ("error", None)
compare_current: str = current_version.lstrip("v") if strip_v_prefix else current_version compare_current: str = current_version.lstrip("v") if strip_v_prefix else current_version
compare_latest: str = latest_tag.lstrip("v") if strip_v_prefix else latest_tag compare_latest: str = latest_tag.lstrip("v") if strip_v_prefix else latest_tag
if self._compare_versions(compare_current, compare_latest): if self._compare_versions(compare_current, compare_latest):
LOG.warning(f"{name} update available: {current_version} -> {latest_tag}") LOG.warning(f"{name} update available: {current_version} -> {latest_tag}")
result: tuple[str, str] = ("update_available", latest_tag) result = ("update_available", latest_tag)
await self._cache.aset(cache_key, result, self.CACHE_DURATION)
return result
LOG.info(f"No {name} updates available.")
result: tuple[str, None] = ("up_to_date", None)
await self._cache.aset(cache_key, result, self.CACHE_DURATION) await self._cache.aset(cache_key, result, self.CACHE_DURATION)
return result return result
LOG.info(f"No {name} updates available.")
result = ("up_to_date", None)
await self._cache.aset(cache_key, result, self.CACHE_DURATION)
return result
except Exception as e: except Exception as e:
LOG.exception(e) LOG.exception(e)
LOG.error(f"Error checking for {name} updates: {e!s}") LOG.error(f"Error checking for {name} updates: {e!s}")

View file

@ -1,17 +1,61 @@
from __future__ import annotations from __future__ import annotations
import asyncio
import functools
import logging import logging
from typing import Any import threading
from dataclasses import dataclass
from typing import Any, Literal, cast, overload
import httpx import httpx
from .cf_solver_shared import is_cf_challenge, solver from .cf_solver_shared import is_cf_challenge, solver
__all__: list[str] = ["async_client", "sync_client"] __all__: list[str] = [
"async_client",
"build_request_headers",
"close_shared_clients",
"get_async_client",
"get_sync_client",
"resolve_curl_transport",
"sync_client",
]
LOG: logging.Logger = logging.getLogger("httpx_cf") LOG: logging.Logger = logging.getLogger("httpx_cf")
class Globals:
random_agent: str | None = None
SHARED_ASYNC_CLIENTS: dict[_AsyncClientKey, httpx.AsyncClient] = {}
SHARED_SYNC_CLIENTS: dict[_SyncClientKey, httpx.Client] = {}
SHARED_CLIENT_LOCK = threading.Lock()
def get_random_agent() -> str:
if Globals.random_agent:
return Globals.random_agent
from yt_dlp.utils.networking import random_user_agent
Globals.random_agent = random_user_agent()
return Globals.random_agent
@dataclass(frozen=True)
class _AsyncClientKey:
enable_cf: bool
proxy: str | None
use_curl: bool
curl_impersonate: str
curl_default_headers: bool
@dataclass(frozen=True)
class _SyncClientKey:
enable_cf: bool
proxy: str | None
def _parse_cookie_header(cookie_header: str | None) -> dict[str, str]: def _parse_cookie_header(cookie_header: str | None) -> dict[str, str]:
cookies: dict[str, str] = {} cookies: dict[str, str] = {}
@ -37,17 +81,95 @@ def _merge_cookies(existing_header: str | None, new_cookies: list[dict[str, str]
return merged return merged
def _normalize_proxy(proxy: str | dict[str, str] | None) -> str | None:
if proxy is None:
return None
if isinstance(proxy, str):
cleaned: str = proxy.strip()
return cleaned or None
if isinstance(proxy, dict):
return "|".join(f"{key}={value}" for key, value in sorted(proxy.items()))
return str(proxy)
@functools.lru_cache(maxsize=1)
def _curl_available() -> bool:
try:
import httpx_curl_cffi # noqa: F401
return True
except Exception:
return False
def resolve_curl_transport(use_curl: bool = True) -> bool:
return use_curl and _curl_available()
def _build_async_curl_transport(
use_curl: bool,
curl_impersonate: str,
curl_default_headers: bool,
) -> httpx.AsyncBaseTransport | None:
if not resolve_curl_transport(use_curl):
return None
from httpx_curl_cffi import AsyncCurlTransport
return AsyncCurlTransport(
impersonate=curl_impersonate,
default_headers=curl_default_headers,
)
@overload
def _get_transport(
enable_cf: bool,
is_async: Literal[True],
transport: httpx.AsyncBaseTransport | None,
) -> httpx.AsyncBaseTransport: ...
@overload
def _get_transport(
enable_cf: bool,
is_async: Literal[False],
transport: httpx.BaseTransport | None,
) -> httpx.BaseTransport: ...
def _get_transport( def _get_transport(
enable_cf: bool, enable_cf: bool,
is_async: bool, is_async: bool,
transport: httpx.AsyncBaseTransport | None, transport: httpx.AsyncBaseTransport | httpx.BaseTransport | None,
) -> httpx.AsyncBaseTransport: ) -> httpx.AsyncBaseTransport | httpx.BaseTransport:
if enable_cf: if enable_cf:
return CFAsyncTransport(base=transport) if is_async else CFTransport(base=transport) if is_async:
async_transport = cast("httpx.AsyncBaseTransport | None", transport)
return CFAsyncTransport(base=async_transport)
sync_transport = cast("httpx.BaseTransport | None", transport)
return CFTransport(base=sync_transport)
return transport or (httpx.AsyncHTTPTransport() if is_async else httpx.HTTPTransport()) return transport or (httpx.AsyncHTTPTransport() if is_async else httpx.HTTPTransport())
def build_request_headers(
base_headers: dict[str, str] | None = None,
user_agent: str | None = None,
use_curl: bool = True,
) -> dict[str, str]:
headers: dict[str, str] = base_headers.copy() if isinstance(base_headers, dict) else {}
if user_agent and not use_curl:
headers.setdefault("User-Agent", user_agent)
return headers
class CFAsyncTransport(httpx.AsyncBaseTransport): class CFAsyncTransport(httpx.AsyncBaseTransport):
def __init__(self, base: httpx.AsyncBaseTransport | None = None): def __init__(self, base: httpx.AsyncBaseTransport | None = None):
self.base: httpx.AsyncBaseTransport | httpx.AsyncHTTPTransport = base or httpx.AsyncHTTPTransport() self.base: httpx.AsyncBaseTransport | httpx.AsyncHTTPTransport = base or httpx.AsyncHTTPTransport()
@ -82,7 +204,9 @@ class CFAsyncTransport(httpx.AsyncBaseTransport):
return await self.base.handle_async_request(request) return await self.base.handle_async_request(request)
def close(self) -> None: def close(self) -> None:
self.base.close() close_fn = getattr(self.base, "close", None)
if callable(close_fn):
close_fn()
class CFTransport(httpx.BaseTransport): class CFTransport(httpx.BaseTransport):
@ -134,8 +258,9 @@ def async_client(enable_cf: bool = True, **kwargs: Any) -> httpx.AsyncClient:
httpx.AsyncClient: The configured httpx.AsyncClient instance. httpx.AsyncClient: The configured httpx.AsyncClient instance.
""" """
transport = kwargs.pop("transport", None) transport = cast("httpx.AsyncBaseTransport | None", kwargs.pop("transport", None))
return httpx.AsyncClient(transport=_get_transport(enable_cf, is_async=True, transport=transport), **kwargs) async_transport = _get_transport(enable_cf, is_async=True, transport=transport)
return httpx.AsyncClient(transport=async_transport, **kwargs)
def sync_client(enable_cf: bool = True, **kwargs: Any) -> httpx.Client: def sync_client(enable_cf: bool = True, **kwargs: Any) -> httpx.Client:
@ -150,5 +275,81 @@ def sync_client(enable_cf: bool = True, **kwargs: Any) -> httpx.Client:
httpx.Client: The configured httpx.Client instance. httpx.Client: The configured httpx.Client instance.
""" """
transport = kwargs.pop("transport", None) transport = cast("httpx.BaseTransport | None", kwargs.pop("transport", None))
return httpx.Client(transport=_get_transport(enable_cf, is_async=False, transport=transport), **kwargs) sync_transport = _get_transport(enable_cf, is_async=False, transport=transport)
return httpx.Client(transport=sync_transport, **kwargs)
def get_async_client(
enable_cf: bool = True,
proxy: str | dict[str, str] | None = None,
use_curl: bool = True,
curl_impersonate: str = "chrome",
curl_default_headers: bool = True,
) -> httpx.AsyncClient:
proxy_key = _normalize_proxy(proxy)
use_curl = resolve_curl_transport(use_curl)
key = _AsyncClientKey(
enable_cf=enable_cf,
proxy=proxy_key,
use_curl=use_curl,
curl_impersonate=curl_impersonate,
curl_default_headers=curl_default_headers,
)
with Globals.SHARED_CLIENT_LOCK:
if key in Globals.SHARED_ASYNC_CLIENTS:
return Globals.SHARED_ASYNC_CLIENTS[key]
transport = _build_async_curl_transport(
use_curl=use_curl,
curl_impersonate=curl_impersonate,
curl_default_headers=curl_default_headers,
)
client = httpx.AsyncClient(
transport=cast(
"httpx.AsyncBaseTransport",
_get_transport(enable_cf, is_async=True, transport=transport),
),
proxy=cast("Any", proxy),
)
Globals.SHARED_ASYNC_CLIENTS[key] = client
return client
def get_sync_client(
enable_cf: bool = True,
proxy: str | dict[str, str] | None = None,
) -> httpx.Client:
proxy_key = _normalize_proxy(proxy)
key = _SyncClientKey(enable_cf=enable_cf, proxy=proxy_key)
with Globals.SHARED_CLIENT_LOCK:
if key in Globals.SHARED_SYNC_CLIENTS:
return Globals.SHARED_SYNC_CLIENTS[key]
client = httpx.Client(
transport=cast(
"httpx.BaseTransport",
_get_transport(enable_cf, is_async=False, transport=None),
),
proxy=cast("Any", proxy),
)
Globals.SHARED_SYNC_CLIENTS[key] = client
return client
async def close_shared_clients(_: Any | None = None) -> None:
async_clients = list(Globals.SHARED_ASYNC_CLIENTS.values())
sync_clients = list(Globals.SHARED_SYNC_CLIENTS.values())
Globals.SHARED_ASYNC_CLIENTS.clear()
Globals.SHARED_SYNC_CLIENTS.clear()
for client in sync_clients:
try:
client.close()
except Exception:
pass
if async_clients:
await asyncio.gather(*(client.aclose() for client in async_clients), return_exceptions=True)

View file

@ -27,6 +27,7 @@ from app.library.downloads import DownloadQueue
from app.library.Events import EventBus, Events from app.library.Events import EventBus, Events
from app.library.HttpAPI import HttpAPI from app.library.HttpAPI import HttpAPI
from app.library.HttpSocket import HttpSocket from app.library.HttpSocket import HttpSocket
from app.library.httpx_client import close_shared_clients
from app.library.Scheduler import Scheduler from app.library.Scheduler import Scheduler
from app.library.Services import Services from app.library.Services import Services
from app.library.sqlite_store import SqliteStore from app.library.sqlite_store import SqliteStore
@ -124,6 +125,7 @@ class Main:
get_task_definitions_repo().attach(self._app) get_task_definitions_repo().attach(self._app)
DownloadQueue.get_instance().attach(self._app) DownloadQueue.get_instance().attach(self._app)
UpdateChecker.get_instance().attach(self._app) UpdateChecker.get_instance().attach(self._app)
self._app.on_shutdown.append(close_shared_clients)
EventBus.get_instance().emit( EventBus.get_instance().emit(
Events.LOADED, Events.LOADED,

View file

@ -1,14 +1,13 @@
import logging import logging
import time import time
from datetime import UTC, datetime from datetime import UTC, datetime
from typing import Any
from aiohttp import web from aiohttp import web
from aiohttp.web import Request, Response from aiohttp.web import Request, Response
from yt_dlp.utils.networking import random_user_agent
from app.library.cache import Cache from app.library.cache import Cache
from app.library.config import Config from app.library.config import Config
from app.library.httpx_client import Globals, build_request_headers, get_async_client, resolve_curl_transport
from app.library.router import add_route, route from app.library.router import add_route, route
from app.library.YTDLPOpts import YTDLPOpts from app.library.YTDLPOpts import YTDLPOpts
@ -65,48 +64,38 @@ async def get_doc(request: Request, config: Config, cache: Cache) -> Response:
try: try:
ytdlp_args: dict = YTDLPOpts.get_instance().preset(name=config.default_preset).get_all() ytdlp_args: dict = YTDLPOpts.get_instance().preset(name=config.default_preset).get_all()
opts: dict[str, Any] = { use_curl = resolve_curl_transport()
request_headers = build_request_headers(
user_agent=request.headers.get("User-Agent", ytdlp_args.get("user_agent", Globals.get_random_agent())),
use_curl=use_curl,
)
proxy = ytdlp_args.get("proxy")
client = get_async_client(proxy=proxy, use_curl=use_curl)
LOG.debug(f"Fetching doc from '{url}'.")
response = await client.request(
method="GET",
url=url,
follow_redirects=True,
headers=request_headers,
)
dct = {
"body": response.content,
"headers": { "headers": {
"User-Agent": request.headers.get("User-Agent", ytdlp_args.get("user_agent", random_user_agent())), "Content-Type": EXT_TO_MIME.get(file[file.rfind(".") :], "text/plain"),
"Pragma": "public",
"Access-Control-Allow-Origin": "*",
"Cache-Control": f"public, max-age={time.time() + 3600}",
"Expires": time.strftime(
"%a, %d %b %Y %H:%M:%S GMT",
datetime.fromtimestamp(time.time() + 3600, tz=UTC).timetuple(),
),
}, },
} }
if proxy := ytdlp_args.get("proxy"):
opts["proxy"] = proxy
try: cache.set(cache_key, dct, ttl=3600)
from httpx_curl_cffi import AsyncCurlTransport, CurlOpt
opts["transport"] = AsyncCurlTransport( return web.Response(**dct)
impersonate="chrome",
default_headers=True,
curl_options={CurlOpt.FRESH_CONNECT: True},
)
opts.pop("headers", None)
except Exception:
pass
from app.library.httpx_client import async_client
async with async_client(**opts) as client:
LOG.debug(f"Fetching doc from '{url}'.")
response = await client.request(method="GET", url=url, follow_redirects=True)
dct = {
"body": response.content,
"headers": {
"Content-Type": EXT_TO_MIME.get(file[file.rfind(".") :], "text/plain"),
"Pragma": "public",
"Access-Control-Allow-Origin": "*",
"Cache-Control": f"public, max-age={time.time() + 3600}",
"Expires": time.strftime(
"%a, %d %b %Y %H:%M:%S GMT",
datetime.fromtimestamp(time.time() + 3600, tz=UTC).timetuple(),
),
},
}
cache.set(cache_key, dct, ttl=3600)
return web.Response(**dct)
except Exception as e: except Exception as e:
LOG.error(f"Failed to request doc from '{url}'.'. '{e!s}'.") LOG.error(f"Failed to request doc from '{url}'.'. '{e!s}'.")
return web.json_response(data={"error": "Failed to get doc."}, status=web.HTTPInternalServerError.status_code) return web.json_response(data={"error": "Failed to get doc."}, status=web.HTTPInternalServerError.status_code)

View file

@ -7,11 +7,11 @@ from urllib.parse import urlparse
from aiohttp import web from aiohttp import web
from aiohttp.web import Request, Response from aiohttp.web import Request, Response
from yt_dlp.utils.networking import random_user_agent
from app.library.ag_utils import ag from app.library.ag_utils import ag
from app.library.cache import Cache from app.library.cache import Cache
from app.library.config import Config from app.library.config import Config
from app.library.httpx_client import Globals, build_request_headers, get_async_client, resolve_curl_transport
from app.library.router import route from app.library.router import route
from app.library.Utils import validate_url from app.library.Utils import validate_url
from app.library.YTDLPOpts import YTDLPOpts from app.library.YTDLPOpts import YTDLPOpts
@ -45,52 +45,39 @@ async def get_thumbnail(request: Request, config: Config) -> Response:
try: try:
ytdlp_args: dict = YTDLPOpts.get_instance().preset(name=config.default_preset).get_all() ytdlp_args: dict = YTDLPOpts.get_instance().preset(name=config.default_preset).get_all()
opts: dict[str, Any] = { use_curl = resolve_curl_transport()
"headers": { request_headers = build_request_headers(
"User-Agent": request.headers.get("User-Agent", ytdlp_args.get("user_agent", random_user_agent())), user_agent=request.headers.get("User-Agent", ytdlp_args.get("user_agent", Globals.get_random_agent())),
use_curl=use_curl,
)
proxy = ytdlp_args.get("proxy")
client = get_async_client(proxy=proxy, use_curl=use_curl)
LOG.debug(f"Fetching thumbnail from '{url}'.")
response = await client.request(
method="GET",
url=url,
follow_redirects=True,
headers=request_headers,
)
if response.status_code != web.HTTPOk.status_code:
LOG.error(f"Failed to fetch thumbnail from '{url}'. Status code: {response.status_code}.")
return web.json_response(data={"error": "failed to retrieve the thumbnail."}, status=response.status_code)
return web.Response(
body=response.content,
headers={
"Content-Type": response.headers.get("Content-Type"),
"Pragma": "public",
"Access-Control-Allow-Origin": "*",
"Cache-Control": f"public, max-age={time.time() + 31536000}",
"Expires": time.strftime(
"%a, %d %b %Y %H:%M:%S GMT",
datetime.fromtimestamp(time.time() + 31536000, tz=UTC).timetuple(),
),
}, },
} )
if proxy := ytdlp_args.get("proxy"):
opts["proxy"] = proxy
try:
from httpx_curl_cffi import AsyncCurlTransport, CurlOpt
opts["transport"] = AsyncCurlTransport(
impersonate="chrome",
default_headers=True,
curl_options={CurlOpt.FRESH_CONNECT: True},
)
opts.pop("headers", None)
except Exception:
pass
from app.library.httpx_client import async_client
async with async_client(**opts) as client:
LOG.debug(f"Fetching thumbnail from '{url}'.")
response = await client.request(method="GET", url=url, follow_redirects=True)
if response.status_code != web.HTTPOk.status_code:
LOG.error(f"Failed to fetch thumbnail from '{url}'. Status code: {response.status_code}.")
return web.json_response(
data={"error": "failed to retrieve the thumbnail."}, status=response.status_code
)
return web.Response(
body=response.content,
headers={
"Content-Type": response.headers.get("Content-Type"),
"Pragma": "public",
"Access-Control-Allow-Origin": "*",
"Cache-Control": f"public, max-age={time.time() + 31536000}",
"Expires": time.strftime(
"%a, %d %b %Y %H:%M:%S GMT",
datetime.fromtimestamp(time.time() + 31536000, tz=UTC).timetuple(),
),
},
)
except Exception as e: except Exception as e:
LOG.error(f"Error fetching thumbnail from '{url}'. '{e}'.") LOG.error(f"Error fetching thumbnail from '{url}'. '{e}'.")
return web.json_response( return web.json_response(
@ -126,95 +113,100 @@ async def get_background(request: Request, config: Config, cache: Cache) -> Resp
CACHE_KEY = "random_background" CACHE_KEY = "random_background"
if cache.has(CACHE_KEY) and not request.query.get("force", False): if cache.has(CACHE_KEY) and not request.query.get("force", False):
data = await cache.aget(CACHE_KEY) cached_data = await cache.aget(CACHE_KEY)
return web.Response( if isinstance(cached_data, dict):
body=data.get("content"), cached_headers = cached_data.get("headers")
headers={ if not isinstance(cached_headers, dict):
"X-Cache": "HIT", cached_headers = {}
"X-Cache-TTL": str(await cache.attl(CACHE_KEY)), cached_headers = {str(key): str(value) for key, value in cached_headers.items()}
"X-Image-Via": data.get("backend"), cached_backend = cached_data.get("backend")
**data.get("headers"), if not isinstance(cached_backend, str):
}, cached_backend = "" if cached_backend is None else str(cached_backend)
) return web.Response(
body=cached_data.get("content"),
ytdlp_args: dict = YTDLPOpts.get_instance().preset(name=config.default_preset).get_all() headers={
opts: dict[str, Any] = { "X-Cache": "HIT",
"headers": { "X-Cache-TTL": str(await cache.attl(CACHE_KEY)),
"User-Agent": request.headers.get("User-Agent", ytdlp_args.get("user_agent", random_user_agent())), "X-Image-Via": cached_backend,
}, **cached_headers,
} },
if proxy := ytdlp_args.get("proxy"):
opts["proxy"] = proxy
try:
from httpx_curl_cffi import AsyncCurlTransport, CurlOpt
opts["transport"] = AsyncCurlTransport(
impersonate="chrome",
default_headers=True,
curl_options={CurlOpt.FRESH_CONNECT: True},
)
opts.pop("headers", None)
except Exception:
pass
from app.library.httpx_client import async_client
async with async_client(**opts) as client:
if backend.startswith("https://www.bing.com/HPImageArchive.aspx"):
if not cache.has(CACHE_KEY_BING):
response = await client.request(method="GET", url=backend)
if response.status_code != web.HTTPOk.status_code:
return web.json_response(
data={"error": "failed to retrieve the random background image."},
status=web.HTTPInternalServerError.status_code,
)
img_url: str | None = ag(response.json(), "images.0.url")
if not img_url:
return web.json_response(
data={"error": "failed to retrieve the random background image."},
status=web.HTTPInternalServerError.status_code,
)
backend = f"https://www.bing.com{img_url}"
await cache.aset(key=CACHE_KEY_BING, value=backend, ttl=3600 * 24)
else:
backend: str = await cache.aget(CACHE_KEY_BING)
LOG.debug(f"Requesting random picture from '{backend!s}'.")
response = await client.request(method="GET", url=backend, follow_redirects=True)
if response.status_code != web.HTTPOk.status_code:
return web.json_response(
data={"error": "failed to retrieve the random background image."},
status=web.HTTPInternalServerError.status_code,
) )
data: dict[str, Any] = { ytdlp_args: dict = YTDLPOpts.get_instance().preset(name=config.default_preset).get_all()
"content": response.content, use_curl = resolve_curl_transport()
"backend": urlparse(backend).netloc, request_headers = build_request_headers(
"headers": { user_agent=request.headers.get("User-Agent", ytdlp_args.get("user_agent", Globals.get_random_agent())),
"Content-Type": response.headers.get("Content-Type", "image/jpeg"), use_curl=use_curl,
"Content-Length": str(len(response.content)), )
}, proxy = ytdlp_args.get("proxy")
}
await cache.aset(key=CACHE_KEY, value=data, ttl=3600) client = get_async_client(proxy=proxy, use_curl=use_curl)
if backend.startswith("https://www.bing.com/HPImageArchive.aspx"):
if not cache.has(CACHE_KEY_BING):
response = await client.request(method="GET", url=backend, headers=request_headers)
if response.status_code != web.HTTPOk.status_code:
return web.json_response(
data={"error": "failed to retrieve the random background image."},
status=web.HTTPInternalServerError.status_code,
)
LOG.debug(f"Random background image from '{backend!s}' cached.") img_url: str | None = ag(response.json(), "images.0.url")
if not img_url:
return web.json_response(
data={"error": "failed to retrieve the random background image."},
status=web.HTTPInternalServerError.status_code,
)
return web.Response( backend = f"https://www.bing.com{img_url}"
body=data.get("content"), await cache.aset(key=CACHE_KEY_BING, value=backend, ttl=3600 * 24)
headers={ else:
"X-Cache": "MISS", backend = await cache.aget(CACHE_KEY_BING)
"X-Cache-TTL": "3600",
"X-Image-Via": data.get("backend"), if not isinstance(backend, str) or not backend:
**data.get("headers"), return web.json_response(
}, data={"error": "failed to retrieve the random background image."},
status=web.HTTPInternalServerError.status_code,
) )
LOG.debug(f"Requesting random picture from '{backend!s}'.")
response = await client.request(
method="GET",
url=backend,
follow_redirects=True,
headers=request_headers,
)
if response.status_code != web.HTTPOk.status_code:
return web.json_response(
data={"error": "failed to retrieve the random background image."},
status=web.HTTPInternalServerError.status_code,
)
data: dict[str, Any] = {
"content": response.content,
"backend": urlparse(backend).netloc,
"headers": {
"Content-Type": response.headers.get("Content-Type", "image/jpeg"),
"Content-Length": str(len(response.content)),
},
}
response_headers: dict[str, str] = data["headers"]
image_via = str(data.get("backend") or "")
await cache.aset(key=CACHE_KEY, value=data, ttl=3600)
LOG.debug(f"Random background image from '{backend!s}' cached.")
return web.Response(
body=data.get("content"),
headers={
"X-Cache": "MISS",
"X-Cache-TTL": "3600",
"X-Image-Via": image_via,
**response_headers,
},
)
except Exception as e: except Exception as e:
LOG.error(f"Failed to request random background image from '{backend!s}'.'. '{e!s}'.") LOG.error(f"Failed to request random background image from '{backend!s}'.'. '{e!s}'.")
return web.json_response( return web.json_response(

View file

@ -153,7 +153,7 @@ class TestUpdateChecker:
assert "" == config.yt_new_version, "Should not update yt_new_version when disabled" assert "" == config.yt_new_version, "Should not update yt_new_version when disabled"
@pytest.mark.asyncio @pytest.mark.asyncio
@patch("app.library.UpdateChecker.async_client") @patch("app.library.UpdateChecker.get_async_client")
async def test_check_for_updates_finds_newer_version(self, mock_client): async def test_check_for_updates_finds_newer_version(self, mock_client):
"""Test that check_for_updates detects when a newer version is available.""" """Test that check_for_updates detects when a newer version is available."""
from app.library.config import Config from app.library.config import Config
@ -173,9 +173,9 @@ class TestUpdateChecker:
mock_ytdlp_response.json.return_value = {"tag_name": "9999.12.31"} mock_ytdlp_response.json.return_value = {"tag_name": "9999.12.31"}
mock_get = AsyncMock(side_effect=[mock_app_response, mock_ytdlp_response]) mock_get = AsyncMock(side_effect=[mock_app_response, mock_ytdlp_response])
mock_context = AsyncMock() mock_http = MagicMock()
mock_context.__aenter__.return_value.get = mock_get mock_http.get = mock_get
mock_client.return_value = mock_context mock_client.return_value = mock_http
with patch("app.library.UpdateChecker.APP_VERSION", "v1.0.0"): with patch("app.library.UpdateChecker.APP_VERSION", "v1.0.0"):
checker = UpdateChecker.get_instance(config=config) checker = UpdateChecker.get_instance(config=config)
@ -190,7 +190,7 @@ class TestUpdateChecker:
assert checker._job_id is None, "Should stop scheduled task after finding app update" assert checker._job_id is None, "Should stop scheduled task after finding app update"
@pytest.mark.asyncio @pytest.mark.asyncio
@patch("app.library.UpdateChecker.async_client") @patch("app.library.UpdateChecker.get_async_client")
async def test_check_for_updates_no_update_available(self, mock_client): async def test_check_for_updates_no_update_available(self, mock_client):
"""Test that check_for_updates correctly handles when no update is available.""" """Test that check_for_updates correctly handles when no update is available."""
from app.library.config import Config from app.library.config import Config
@ -210,9 +210,9 @@ class TestUpdateChecker:
mock_ytdlp_response.json.return_value = {"tag_name": "2020.01.01"} mock_ytdlp_response.json.return_value = {"tag_name": "2020.01.01"}
mock_get = AsyncMock(side_effect=[mock_app_response, mock_ytdlp_response]) mock_get = AsyncMock(side_effect=[mock_app_response, mock_ytdlp_response])
mock_context = AsyncMock() mock_http = MagicMock()
mock_context.__aenter__.return_value.get = mock_get mock_http.get = mock_get
mock_client.return_value = mock_context mock_client.return_value = mock_http
checker = UpdateChecker.get_instance(config=config) checker = UpdateChecker.get_instance(config=config)
checker._job_id = "test-job" checker._job_id = "test-job"
@ -228,7 +228,7 @@ class TestUpdateChecker:
assert "test-job" == checker._job_id, "Should keep scheduled task running" assert "test-job" == checker._job_id, "Should keep scheduled task running"
@pytest.mark.asyncio @pytest.mark.asyncio
@patch("app.library.UpdateChecker.async_client") @patch("app.library.UpdateChecker.get_async_client")
async def test_check_for_updates_handles_http_error(self, mock_client): async def test_check_for_updates_handles_http_error(self, mock_client):
"""Test that check_for_updates handles HTTP errors gracefully.""" """Test that check_for_updates handles HTTP errors gracefully."""
from app.library.config import Config from app.library.config import Config
@ -242,9 +242,9 @@ class TestUpdateChecker:
mock_response = MagicMock() mock_response = MagicMock()
mock_response.status_code = 404 mock_response.status_code = 404
mock_context = AsyncMock() mock_http = MagicMock()
mock_context.__aenter__.return_value.get = AsyncMock(return_value=mock_response) mock_http.get = AsyncMock(return_value=mock_response)
mock_client.return_value = mock_context mock_client.return_value = mock_http
checker = UpdateChecker.get_instance(config=config) checker = UpdateChecker.get_instance(config=config)
@ -258,7 +258,7 @@ class TestUpdateChecker:
assert "" == config.yt_new_version, "Should not set yt_new_version on HTTP error" assert "" == config.yt_new_version, "Should not set yt_new_version on HTTP error"
@pytest.mark.asyncio @pytest.mark.asyncio
@patch("app.library.UpdateChecker.async_client") @patch("app.library.UpdateChecker.get_async_client")
async def test_check_for_updates_handles_exception(self, mock_client): async def test_check_for_updates_handles_exception(self, mock_client):
"""Test that check_for_updates handles exceptions gracefully.""" """Test that check_for_updates handles exceptions gracefully."""
from app.library.config import Config from app.library.config import Config
@ -269,9 +269,9 @@ class TestUpdateChecker:
config.new_version = "" config.new_version = ""
config.yt_new_version = "" config.yt_new_version = ""
mock_context = AsyncMock() mock_http = MagicMock()
mock_context.__aenter__.return_value.get = AsyncMock(side_effect=Exception("Network error")) mock_http.get = AsyncMock(side_effect=Exception("Network error"))
mock_client.return_value = mock_context mock_client.return_value = mock_http
checker = UpdateChecker.get_instance(config=config) checker = UpdateChecker.get_instance(config=config)
@ -337,10 +337,10 @@ class TestUpdateChecker:
mock_ytdlp_response.json.return_value = {"tag_name": "2026.01.01"} mock_ytdlp_response.json.return_value = {"tag_name": "2026.01.01"}
mock_get = AsyncMock(side_effect=[mock_app_response, mock_ytdlp_response]) mock_get = AsyncMock(side_effect=[mock_app_response, mock_ytdlp_response])
mock_context = AsyncMock() mock_http = MagicMock()
mock_context.__aenter__.return_value.get = mock_get mock_http.get = mock_get
with patch("app.library.UpdateChecker.async_client", return_value=mock_context): with patch("app.library.UpdateChecker.get_async_client", return_value=mock_http):
with patch("app.library.UpdateChecker.APP_VERSION", "v1.0.0"): with patch("app.library.UpdateChecker.APP_VERSION", "v1.0.0"):
checker = UpdateChecker.get_instance(config=config) checker = UpdateChecker.get_instance(config=config)
await checker.check_for_updates() await checker.check_for_updates()
@ -383,7 +383,7 @@ class TestUpdateChecker:
loop.close() loop.close()
@pytest.mark.asyncio @pytest.mark.asyncio
@patch("app.library.UpdateChecker.async_client") @patch("app.library.UpdateChecker.get_async_client")
async def test_check_ytdlp_version_finds_newer_version(self, mock_client): async def test_check_ytdlp_version_finds_newer_version(self, mock_client):
"""Test that yt-dlp check detects when a newer version is available.""" """Test that yt-dlp check detects when a newer version is available."""
from app.library.config import Config from app.library.config import Config
@ -397,9 +397,9 @@ class TestUpdateChecker:
mock_response.status_code = 200 mock_response.status_code = 200
mock_response.json.return_value = {"tag_name": "9999.12.31"} mock_response.json.return_value = {"tag_name": "9999.12.31"}
mock_context = AsyncMock() mock_http = MagicMock()
mock_context.__aenter__.return_value.get = AsyncMock(return_value=mock_response) mock_http.get = AsyncMock(return_value=mock_response)
mock_client.return_value = mock_context mock_client.return_value = mock_http
checker = UpdateChecker.get_instance(config=config) checker = UpdateChecker.get_instance(config=config)
status, new_version = await checker._check_ytdlp_version() status, new_version = await checker._check_ytdlp_version()
@ -409,7 +409,7 @@ class TestUpdateChecker:
assert "9999.12.31" == config.yt_new_version, "Should store new yt-dlp version tag" assert "9999.12.31" == config.yt_new_version, "Should store new yt-dlp version tag"
@pytest.mark.asyncio @pytest.mark.asyncio
@patch("app.library.UpdateChecker.async_client") @patch("app.library.UpdateChecker.get_async_client")
async def test_check_ytdlp_version_no_update_available(self, mock_client): async def test_check_ytdlp_version_no_update_available(self, mock_client):
"""Test that yt-dlp check correctly handles when no update is available.""" """Test that yt-dlp check correctly handles when no update is available."""
from app.library.config import Config from app.library.config import Config
@ -423,9 +423,9 @@ class TestUpdateChecker:
mock_response.status_code = 200 mock_response.status_code = 200
mock_response.json.return_value = {"tag_name": "2020.01.01"} mock_response.json.return_value = {"tag_name": "2020.01.01"}
mock_context = AsyncMock() mock_http = MagicMock()
mock_context.__aenter__.return_value.get = AsyncMock(return_value=mock_response) mock_http.get = AsyncMock(return_value=mock_response)
mock_client.return_value = mock_context mock_client.return_value = mock_http
checker = UpdateChecker.get_instance(config=config) checker = UpdateChecker.get_instance(config=config)
status, new_version = await checker._check_ytdlp_version() status, new_version = await checker._check_ytdlp_version()
@ -435,7 +435,7 @@ class TestUpdateChecker:
assert "" == config.yt_new_version, "Should clear yt_new_version when no update available" assert "" == config.yt_new_version, "Should clear yt_new_version when no update available"
@pytest.mark.asyncio @pytest.mark.asyncio
@patch("app.library.UpdateChecker.async_client") @patch("app.library.UpdateChecker.get_async_client")
async def test_check_ytdlp_version_handles_http_error(self, mock_client): async def test_check_ytdlp_version_handles_http_error(self, mock_client):
"""Test that yt-dlp check handles HTTP errors gracefully.""" """Test that yt-dlp check handles HTTP errors gracefully."""
from app.library.config import Config from app.library.config import Config
@ -448,9 +448,9 @@ class TestUpdateChecker:
mock_response = MagicMock() mock_response = MagicMock()
mock_response.status_code = 500 mock_response.status_code = 500
mock_context = AsyncMock() mock_http = MagicMock()
mock_context.__aenter__.return_value.get = AsyncMock(return_value=mock_response) mock_http.get = AsyncMock(return_value=mock_response)
mock_client.return_value = mock_context mock_client.return_value = mock_http
checker = UpdateChecker.get_instance(config=config) checker = UpdateChecker.get_instance(config=config)
status, new_version = await checker._check_ytdlp_version() status, new_version = await checker._check_ytdlp_version()
@ -460,7 +460,7 @@ class TestUpdateChecker:
assert "" == config.yt_new_version, "Should not set yt_new_version on HTTP error" assert "" == config.yt_new_version, "Should not set yt_new_version on HTTP error"
@pytest.mark.asyncio @pytest.mark.asyncio
@patch("app.library.UpdateChecker.async_client") @patch("app.library.UpdateChecker.get_async_client")
async def test_check_for_updates_caches_separately(self, mock_client): async def test_check_for_updates_caches_separately(self, mock_client):
"""Test that app and yt-dlp checks are cached separately.""" """Test that app and yt-dlp checks are cached separately."""
from app.library.config import Config from app.library.config import Config
@ -480,9 +480,9 @@ class TestUpdateChecker:
mock_ytdlp_response.json.return_value = {"tag_name": "2026.01.01"} mock_ytdlp_response.json.return_value = {"tag_name": "2026.01.01"}
mock_get = AsyncMock(side_effect=[mock_app_response, mock_ytdlp_response]) mock_get = AsyncMock(side_effect=[mock_app_response, mock_ytdlp_response])
mock_context = AsyncMock() mock_http = MagicMock()
mock_context.__aenter__.return_value.get = mock_get mock_http.get = mock_get
mock_client.return_value = mock_context mock_client.return_value = mock_http
with patch("app.library.UpdateChecker.APP_VERSION", "v1.0.0"): with patch("app.library.UpdateChecker.APP_VERSION", "v1.0.0"):
checker = UpdateChecker.get_instance(config=config) checker = UpdateChecker.get_instance(config=config)