Files

412 lines
12 KiB
Python
Raw Permalink Normal View History

import asyncio as aio
import mimetypes
import random
import sys
import time
from abc import ABC, abstractmethod
from collections.abc import AsyncIterable, Callable
from math import floor
from pathlib import Path
from typing import Generic, NamedTuple, ParamSpec, TypeAlias, TypedDict, TypeVar
from typing_extensions import override
from cookit.common import race
from cookit.loguru import warning_suppress
from httpx import AsyncClient, Response
from nonebot import get_driver, logger
from .config import BG_PRELOAD_CACHE_DIR, DEFAULT_BG_PATH, config
if sys.version_info >= (3, 11):
from asyncio.taskgroups import TaskGroup
else:
from taskgroup import TaskGroup
class BgBytesData(NamedTuple):
data: bytes | None
mime: str
class BgFileData(NamedTuple):
path: Path | None
mime: str
BgData: TypeAlias = BgBytesData | BgFileData
BGProviderType = Callable[[int], AsyncIterable[BgData]]
T = TypeVar("T")
TBP = TypeVar("TBP", bound=BGProviderType)
P = ParamSpec("P")
DEFAULT_MIME = "application/octet-stream"
registered_bg_providers: dict[str, BGProviderType] = {}
def get_bg_files() -> list["Path"]:
if not config.ps_bg_local_path.exists():
logger.warning("Custom background path does not exist, fallback to default")
return [DEFAULT_BG_PATH]
if config.ps_bg_local_path.is_file():
return [config.ps_bg_local_path]
files = [x for x in config.ps_bg_local_path.glob("*") if x.is_file()]
if not files:
logger.warning("Custom background dir has no file in it, fallback to default")
return [DEFAULT_BG_PATH]
return files
BG_FILES = get_bg_files()
def refresh_bg_files():
global BG_FILES
BG_FILES = get_bg_files()
def bg_provider(name: str | None = None):
def deco(func: TBP) -> TBP:
provider_name = name or func.__name__
if provider_name in registered_bg_providers:
raise ValueError(f"Duplicate bg provider name `{provider_name}`")
registered_bg_providers[provider_name] = func
return func
return deco
def iter_batch_sizes(size: int, max_size: int):
if size <= max_size:
yield size
else:
full_sizes = floor(max_size / size)
for _ in range(full_sizes):
yield max_size
if rest_count := full_sizes * max_size:
yield rest_count
def resp_to_bg_data(resp: Response):
return BgBytesData(
resp.content,
(resp.headers.get("Content-Type") or DEFAULT_MIME),
)
class CoIterator(ABC, Generic[T]):
def __init__(self):
self.queue = aio.Queue[T | None]()
@abstractmethod
async def run_tasks(self): ...
async def run(self):
await self.run_tasks()
await self.queue.put(None)
async def __aiter__(self):
async with TaskGroup() as t:
t.create_task(self.run())
while (x := await self.queue.get()) is not None:
yield x
@bg_provider("loli")
class LoliBGProvider(CoIterator[BgData]):
def __init__(self, num: int):
super().__init__()
self.num = num
self.sem = aio.Semaphore(4)
async def task_piece(self, cli: AsyncClient):
async with self.sem:
with warning_suppress("Failed to fetch image"):
x = resp_to_bg_data(
(
await cli.get("https://www.loliapi.com/acg/pe/")
).raise_for_status(),
)
await self.queue.put(x)
@override
async def run_tasks(self):
async with AsyncClient(
follow_redirects=True,
proxy=config.proxy,
timeout=config.ps_req_timeout,
) as cli:
await aio.gather(*(self.task_piece(cli) for _ in range(self.num)))
class LoliconRespDataUrls(TypedDict):
original: str
class LoliconRespData(TypedDict):
urls: LoliconRespDataUrls
class LoliconResp(TypedDict):
data: list[LoliconRespData]
@bg_provider("lolicon")
class LoliconBGProvider(CoIterator[BgData]):
def __init__(self, num: int):
super().__init__()
self.num = num
self.sem = aio.Semaphore(4)
self.url_queue = aio.Queue[str | None]()
async def do_fetch_urls_piece(self, num: int, cli: AsyncClient):
with warning_suppress("Failed to fetch urls"):
resp = await cli.get(
"https://api.lolicon.app/setu/v2",
params={
"num": num,
"r18": config.ps_bg_lolicon_r18_type,
"proxy": "false",
"excludeAI": "true",
},
)
data: LoliconResp = resp.raise_for_status().json()
for x in data["data"]:
await self.url_queue.put(x["urls"]["original"])
async def fetch_urls_task_f(self):
async with AsyncClient(
follow_redirects=True,
proxy=config.proxy,
timeout=config.ps_req_timeout,
) as cli:
for x in iter_batch_sizes(self.num, 20):
await self.do_fetch_urls_piece(x, cli)
await self.url_queue.put(None)
async def fetch_image(self, url: str, cli: AsyncClient):
async with self.sem:
with warning_suppress("Failed to fetch image"):
bg = resp_to_bg_data((await cli.get(url)).raise_for_status())
await self.queue.put(bg)
@override
async def run_tasks(self):
pixiv_client = AsyncClient(
follow_redirects=True,
proxy=config.proxy,
timeout=config.ps_req_timeout,
headers={
"User-Agent": (
"Mozilla/5.0 (Windows NT 10.0; Win64; x64) "
"AppleWebKit/537.36 (KHTML, like Gecko) "
"Chrome/119.0.0.0 "
"Safari/537.36"
),
"Referer": "https://www.pixiv.net/",
},
)
async with TaskGroup() as t, pixiv_client:
t.create_task(self.fetch_urls_task_f())
while (x := await self.url_queue.get()) is not None:
t.create_task(self.fetch_image(x, pixiv_client))
@bg_provider()
async def local(num: int):
files = random.sample(BG_FILES, num)
# logger.debug(f"Chosen background `{files}`")
for x in files:
yield BgFileData(
x,
mimetypes.guess_type(x)[0] or DEFAULT_MIME,
)
def create_none_bg():
return BgBytesData(None, DEFAULT_MIME)
@bg_provider()
async def none(num: int):
for _ in range(num):
yield create_none_bg()
async def fetch_bg(num: int) -> AsyncIterable[BgData]:
if config.ps_bg_provider not in registered_bg_providers:
logger.warning(
f"Unknown background provider `{config.ps_bg_provider}`, fallback to local",
)
async for x in local(num):
yield x
return
# at least we should return one image (x)
# has_img = False
try:
provider = registered_bg_providers[config.ps_bg_provider]
async for x in provider(num):
# has_img = True
yield x
except Exception:
logger.exception(
"Error when getting background, fallback to get one local bg",
)
async for x in local(1):
yield x
# else:
# if has_img:
# return
# logger.warning(
# "Background provider returned empty iterator, fallback to get one local bg",
# )
# async for x in local(1):
# yield x
def cache_bg(bg: BgBytesData):
if not bg.data:
return BgFileData(None, bg.mime)
BG_PRELOAD_CACHE_DIR.mkdir(parents=True, exist_ok=True)
path = BG_PRELOAD_CACHE_DIR / f"{time.time_ns()}.{bg.mime.split('/')[-1]}"
path.write_bytes(bg.data)
return BgFileData(path, bg.mime)
def read_cached_bg_file(bg: BgFileData) -> BgBytesData | None:
if not bg.path:
return BgBytesData(None, bg.mime)
with warning_suppress("Failed to read cached file"):
data = bg.path.read_bytes()
if bg.path.is_relative_to(BG_PRELOAD_CACHE_DIR):
with warning_suppress("Failed to unlink cached file"):
bg.path.unlink()
return BgBytesData(data, bg.mime)
return None
async def get_one_fallback() -> BgBytesData:
with warning_suppress("Failed to get local bg file, fallback to none"):
async for x in local(1):
if bg := read_cached_bg_file(x):
return bg
logger.warning("Failed to read local bg file, fallback to none")
return create_none_bg()
class BgPreloader:
def __init__(self, preload_count: int):
# if preload_count < 1:
# raise ValueError("preload_count must be greater than or equals 1")
self.preload_count = preload_count
self.background_queue = aio.Queue[BgData]()
self.current_load_task_main: aio.Task | None = None
self.consumed_in_loading: bool = False
self.image_got_signal = aio.Event()
self.fire_tasks: set[aio.Task] = set()
# we allow fetch_bg return less image than we require
async def preload_task(
self,
count: int,
fire: bool = False,
fire_done_signal: aio.Event | None = None,
):
logger.debug(f"Preload task started, will preload {count} images, {fire=}")
try:
async for x in fetch_bg(count):
logger.debug("Got one image")
if self.preload_count > 0 or (
fire_done_signal and fire_done_signal.is_set()
):
x = cache_bg(x) if isinstance(x, BgBytesData) else x
await self.background_queue.put(x)
self.image_got_signal.set()
self.image_got_signal.clear()
except Exception:
logger.exception("Unexpected error occurred in preload task")
else:
logger.debug("Preload task finished")
if fire:
return
if (
self.consumed_in_loading
or self.background_queue.qsize() < self.preload_count
):
self.consumed_in_loading = False
self.start_preload()
else:
self.current_load_task_main = None
def start_preload(self, force: bool = False):
count = self.preload_count - self.background_queue.qsize()
if count <= 0 and not force:
logger.debug(
"Current background queue size meets preload count, skip preload",
)
return
task = aio.create_task(self.preload_task(count))
self.current_load_task_main = task
def set_defer_preload(self):
if self.current_load_task_main:
logger.debug("Main preload task already running, set flag")
self.consumed_in_loading = True
else:
self.start_preload()
async def _get_on_fire(self) -> BgBytesData:
task_done_signal = aio.Event()
fire_task = aio.create_task(
self.preload_task(1, fire=True, fire_done_signal=task_done_signal),
)
fire_task.add_done_callback(lambda _: task_done_signal.set())
fire_task.add_done_callback(lambda _: self.fire_tasks.discard(fire_task))
self.fire_tasks.add(fire_task)
try:
await race(
# self.image_got_signal.wait(), # lazy to handle this racing condition now
task_done_signal.wait(),
aio.sleep(15),
)
finally:
task_done_signal.set()
# fire_task.cancel() # should we cancel here? i'm letting it cache to queue
if not self.background_queue.empty():
bg = await self.background_queue.get()
self.set_defer_preload()
if (not isinstance(bg, BgFileData)) or (bg := read_cached_bg_file(bg)):
return bg
logger.error("Unable to get an background image, falling back to local")
return await get_one_fallback()
async def get(self) -> BgBytesData:
self.set_defer_preload()
while not self.background_queue.empty():
bg = await self.background_queue.get()
self.set_defer_preload()
if (not isinstance(bg, BgFileData)) or (bg := read_cached_bg_file(bg)):
return bg
# normally all items in queue should be valid
# if they not, we should fetch
return await self._get_on_fire()
bg_preloader = BgPreloader(config.ps_bg_preload_count)
driver = get_driver()
@driver.on_shutdown
async def _():
for t in bg_preloader.fire_tasks:
t.cancel()