| 1 | from __future__ import annotations |
| 2 | |
| 3 | from abc import ABC |
| 4 | from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set |
| 5 | |
| 6 | from core.downloader_base import DownloadResult |
| 7 | from utils.logger import setup_logger |
| 8 | |
| 9 | if TYPE_CHECKING: |
| 10 | from core.user_downloader import UserDownloader |
| 11 | |
| 12 | logger = setup_logger("UserModeStrategy") |
| 13 | |
| 14 | _MEDIA_TYPE_CHOICES = {"video", "gallery"} |
| 15 | |
| 16 | |
| 17 | class BaseUserModeStrategy(ABC): |
| 18 | mode_name = "" |
| 19 | api_method_name = "" |
| 20 | |
| 21 | def __init__(self, downloader: "UserDownloader"): |
| 22 | self.downloader = downloader |
| 23 | |
| 24 | async def download_mode( |
| 25 | self, |
| 26 | sec_uid: str, |
| 27 | user_info: Dict[str, Any], |
| 28 | seen_aweme_ids: Optional[set[str]] = None, |
| 29 | ) -> DownloadResult: |
| 30 | items = await self.collect_items(sec_uid, user_info) |
| 31 | items = self.apply_filters(items) |
| 32 | author_name = user_info.get("nickname", "unknown") |
| 33 | if seen_aweme_ids is None: |
| 34 | seen_aweme_ids = set() |
| 35 | return await self.downloader._download_mode_items( |
| 36 | mode=self.mode_name, |
| 37 | items=items, |
| 38 | author_name=author_name, |
| 39 | seen_aweme_ids=seen_aweme_ids, |
| 40 | ) |
| 41 | |
| 42 | async def collect_items(self, sec_uid: str, user_info: Dict[str, Any]) -> List[Dict[str, Any]]: |
| 43 | return await self._collect_paged_aweme(sec_uid, user_info) |
| 44 | |
| 45 | def apply_filters(self, items: List[Dict[str, Any]]) -> List[Dict[str, Any]]: |
| 46 | filtered = self._filter_pinned_items(items) |
| 47 | filtered = self.downloader._filter_by_time(filtered) |
| 48 | filtered = self._filter_by_media_type(filtered) |
| 49 | return self.downloader._limit_count(filtered, self.mode_name) |
| 50 | |
| 51 | def _filter_pinned_items(self, items: List[Dict[str, Any]]) -> List[Dict[str, Any]]: |
| 52 | filterer = getattr(self.downloader, "_filter_pinned_items", None) |
| 53 | if callable(filterer): |
| 54 | return filterer(items) |
| 55 | return items |
| 56 | |
| 57 | def _configured_media_types(self) -> Optional[Set[str]]: |
| 58 | if self.mode_name == "music": |
| 59 | return None |
| 60 | raw = self.downloader.config.get("media_types", None) |
| 61 | if not isinstance(raw, (list, tuple, set)): |
| 62 | return None |
| 63 | media_types = {value for value in raw if isinstance(value, str)} |
| 64 | selected = media_types.intersection(_MEDIA_TYPE_CHOICES) |
| 65 | if not selected or selected == _MEDIA_TYPE_CHOICES: |
| 66 | return None |
| 67 | return selected |
| 68 | |
| 69 | def _media_type_filter_enabled(self) -> bool: |
| 70 | return self._configured_media_types() is not None |
| 71 | |
| 72 | def _filter_by_media_type(self, items: List[Dict[str, Any]]) -> List[Dict[str, Any]]: |
| 73 | selected = self._configured_media_types() |
| 74 | if selected is None: |
| 75 | return items |
| 76 | detector = getattr(self.downloader, "_detect_media_type", None) |
| 77 | if not callable(detector): |
| 78 | return items |
| 79 | return [item for item in items if detector(item) in selected] |
| 80 | |
| 81 | async def _collect_paged_aweme( |
| 82 | self, sec_uid: str, user_info: Dict[str, Any] |
| 83 | ) -> List[Dict[str, Any]]: |
| 84 | fetcher = getattr(self.downloader.api_client, self.api_method_name, None) |
| 85 | if not callable(fetcher): |
| 86 | logger.warning( |
| 87 | "Mode %s skipped: API method %s not implemented", |
| 88 | self.mode_name, |
| 89 | self.api_method_name, |
| 90 | ) |
| 91 | return [] |
| 92 | |
| 93 | aweme_list: List[Dict[str, Any]] = [] |
| 94 | max_cursor = 0 |
| 95 | has_more = True |
| 96 | |
| 97 | number_limit = int(self.downloader.config.get("number", {}).get(self.mode_name, 0) or 0) |
| 98 | media_filter_enabled = self._media_type_filter_enabled() |
| 99 | increase_enabled = bool( |
| 100 | self.downloader.config.get("increase", {}).get(self.mode_name, False) |
| 101 | ) |
| 102 | stop_at_downloaded_aweme = ( |
| 103 | increase_enabled and self.mode_name == "like" and self.downloader.database |
| 104 | ) |
| 105 | latest_time = None |
| 106 | if increase_enabled and self.downloader.database and not stop_at_downloaded_aweme: |
| 107 | latest_time = await self.downloader.database.get_latest_aweme_time(user_info.get("uid")) |
| 108 | |
| 109 | while has_more: |
| 110 | await self.downloader.rate_limiter.acquire() |
| 111 | request_cursor = max_cursor |
| 112 | page_data = await fetcher(sec_uid, request_cursor, 20) |
| 113 | page = self._normalize_page_data(page_data) |
| 114 | page_items = self.select_items(page) |
| 115 | if not page_items: |
| 116 | break |
| 117 | |
| 118 | if stop_at_downloaded_aweme: |
| 119 | new_items = [] |
| 120 | for item in page_items: |
| 121 | if await self._is_downloaded_aweme(item): |
| 122 | break |
| 123 | new_items.append(item) |
| 124 | aweme_list.extend(new_items) |
| 125 | if len(new_items) < len(page_items): |
| 126 | break |
| 127 | elif increase_enabled and latest_time: |
| 128 | new_items = [a for a in page_items if a.get("create_time", 0) > latest_time] |
| 129 | aweme_list.extend(new_items) |
| 130 | if len(new_items) < len(page_items): |
| 131 | break |
| 132 | else: |
| 133 | aweme_list.extend(page_items) |
| 134 | |
| 135 | if number_limit > 0: |
| 136 | if media_filter_enabled: |
| 137 | if len(self._filter_by_media_type(aweme_list)) >= number_limit: |
| 138 | break |
| 139 | elif len(aweme_list) >= number_limit: |
| 140 | aweme_list = aweme_list[:number_limit] |
| 141 | break |
| 142 | |
| 143 | has_more = bool(page.get("has_more", False)) |
| 144 | max_cursor = int(page.get("max_cursor", 0) or 0) |
| 145 | if has_more and max_cursor == request_cursor: |
| 146 | logger.warning( |
| 147 | "Mode %s cursor did not advance (%s), stop paging", |
| 148 | self.mode_name, |
| 149 | max_cursor, |
| 150 | ) |
| 151 | break |
| 152 | |
| 153 | return aweme_list |
| 154 | |
| 155 | async def _is_downloaded_aweme(self, item: Dict[str, Any]) -> bool: |
| 156 | aweme_id = str(item.get("aweme_id") or "").strip() |
| 157 | if not aweme_id or not self.downloader.database: |
| 158 | return False |
| 159 | return await self.downloader.database.is_downloaded(aweme_id) |
| 160 | |
| 161 | def select_items(self, page_data: Dict[str, Any]) -> List[Dict[str, Any]]: |
| 162 | items = page_data.get("items") |
| 163 | if isinstance(items, list): |
| 164 | return [item for item in items if isinstance(item, dict)] |
| 165 | return [] |
| 166 | |
| 167 | async def _collect_paged_entries( |
| 168 | self, |
| 169 | fetcher, |
| 170 | *fetch_args: Any, |
| 171 | count: int = 20, |
| 172 | ) -> List[Dict[str, Any]]: |
| 173 | entries: List[Dict[str, Any]] = [] |
| 174 | max_cursor = 0 |
| 175 | has_more = True |
| 176 | |
| 177 | while has_more: |
| 178 | await self.downloader.rate_limiter.acquire() |
| 179 | request_cursor = max_cursor |
| 180 | page_data = await fetcher(*fetch_args, request_cursor, count) |
| 181 | page = self._normalize_page_data(page_data) |
| 182 | page_items = self.select_items(page) |
| 183 | if not page_items: |
| 184 | break |
| 185 | |
| 186 | entries.extend(page_items) |
| 187 | has_more = bool(page.get("has_more", False)) |
| 188 | max_cursor = int(page.get("max_cursor", 0) or 0) |
| 189 | if has_more and max_cursor == request_cursor: |
| 190 | logger.warning( |
| 191 | "Mode %s cursor did not advance (%s), stop paging", |
| 192 | self.mode_name, |
| 193 | max_cursor, |
| 194 | ) |
| 195 | break |
| 196 | |
| 197 | return entries |
| 198 | |
| 199 | async def _expand_metadata_items( |
| 200 | self, |
| 201 | raw_items: List[Dict[str, Any]], |
| 202 | id_field: str, |
| 203 | id_aliases: List[str], |
| 204 | fetch_method_name: str, |
| 205 | ) -> List[Dict[str, Any]]: |
| 206 | """Shared expansion logic for mix/music strategies that receive metadata |
| 207 | items instead of aweme items. Fetches the actual aweme list for each |
| 208 | metadata entry using the given API method.""" |
| 209 | fetcher = getattr(self.downloader.api_client, fetch_method_name, None) |
| 210 | if not callable(fetcher): |
| 211 | return [] |
| 212 | |
| 213 | expanded: List[Dict[str, Any]] = [] |
| 214 | seen_aweme: set[str] = set() |
| 215 | |
| 216 | for item in raw_items: |
| 217 | entry_id = item.get(id_field) |
| 218 | if not entry_id: |
| 219 | for alias in id_aliases: |
| 220 | candidate = item.get(alias) |
| 221 | if not candidate: |
| 222 | info = item.get(f"{id_field.split('_')[0]}_info") |
| 223 | if isinstance(info, dict): |
| 224 | candidate = info.get(id_field) or info.get("id") |
| 225 | if candidate: |
| 226 | entry_id = candidate |
| 227 | break |
| 228 | if not entry_id: |
| 229 | continue |
| 230 | |
| 231 | cursor = 0 |
| 232 | has_more = True |
| 233 | while has_more: |
| 234 | await self.downloader.rate_limiter.acquire() |
| 235 | try: |
| 236 | page_data = await fetcher(str(entry_id), cursor=cursor, count=20) |
| 237 | except Exception as exc: |
| 238 | logger.warning( |
| 239 | "Expansion fetch failed for %s=%s: %s", |
| 240 | id_field, |
| 241 | entry_id, |
| 242 | exc, |
| 243 | ) |
| 244 | break |
| 245 | page = self._normalize_page_data(page_data) |
| 246 | page_items = page.get("items", []) |
| 247 | if not page_items: |
| 248 | break |
| 249 | |
| 250 | for aweme in page_items: |
| 251 | extracted = self._extract_aweme_from_item(aweme) |
| 252 | if not extracted: |
| 253 | continue |
| 254 | aweme_id = str(extracted.get("aweme_id") or "") |
| 255 | if not aweme_id or aweme_id in seen_aweme: |
| 256 | continue |
| 257 | seen_aweme.add(aweme_id) |
| 258 | expanded.append(extracted) |
| 259 | |
| 260 | has_more = bool(page.get("has_more", False)) |
| 261 | next_cursor = int(page.get("max_cursor", 0) or 0) |
| 262 | if has_more and next_cursor == cursor: |
| 263 | logger.warning( |
| 264 | "%s %s cursor did not advance", |
| 265 | id_field, |
| 266 | entry_id, |
| 267 | ) |
| 268 | break |
| 269 | cursor = next_cursor |
| 270 | |
| 271 | return expanded |
| 272 | |
| 273 | @staticmethod |
| 274 | def _extract_aweme_from_item(item: Any) -> Optional[Dict[str, Any]]: |
| 275 | if not isinstance(item, dict): |
| 276 | return None |
| 277 | if item.get("aweme_id"): |
| 278 | return item |
| 279 | for key in ("aweme", "aweme_info", "aweme_detail"): |
| 280 | value = item.get(key) |
| 281 | if isinstance(value, dict) and value.get("aweme_id"): |
| 282 | return value |
| 283 | return None |
| 284 | |
| 285 | @staticmethod |
| 286 | def _normalize_page_data(data: Any) -> Dict[str, Any]: |
| 287 | if not isinstance(data, dict): |
| 288 | return {"items": [], "has_more": False, "max_cursor": 0, "status_code": -1} |
| 289 | |
| 290 | if isinstance(data.get("items"), list): |
| 291 | return { |
| 292 | "items": data.get("items") or [], |
| 293 | "has_more": bool(data.get("has_more")), |
| 294 | "max_cursor": int(data.get("max_cursor", 0) or 0), |
| 295 | "status_code": int(data.get("status_code", 0) or 0), |
| 296 | "raw": data.get("raw", data), |
| 297 | "risk_flags": data.get("risk_flags", {}), |
| 298 | } |
| 299 | |
| 300 | raw_items = data.get("aweme_list") or [] |
| 301 | return { |
| 302 | "items": raw_items if isinstance(raw_items, list) else [], |
| 303 | "has_more": bool(data.get("has_more")), |
| 304 | "max_cursor": int(data.get("max_cursor", 0) or 0), |
| 305 | "status_code": int(data.get("status_code", 0) or 0), |
| 306 | "raw": data, |
| 307 | "risk_flags": {}, |
| 308 | } |
| 309 |