返回 douyin-downloader
base_strategy.py
根目录 / core / user_modes / base_strategy.py
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
309 lines PYTHON