chore: remove unused imports and fix function name conflicts (#5764)

- Remove unused imports in anthropic.py, tmdbv3api/__init__.py, tv.py, test files
- Rename conflicting function names in subscribe.py and webhook.py
- Clean up unused re-exports in tmdbv3api/__init__.py (15 unused exports)
- Apply consistent formatting across API endpoints
This commit is contained in:
DDSRem
2026-05-13 18:59:03 +08:00
committed by GitHub
parent fcf6e14ac9
commit 5a585839ba
34 changed files with 2974 additions and 1838 deletions
+10 -4
View File
@@ -1,6 +1,5 @@
import asyncio import asyncio
import json import json
import time
import uuid import uuid
from typing import AsyncIterator, List, Optional from typing import AsyncIterator, List, Optional
@@ -11,9 +10,12 @@ from app import schemas
from app.api.endpoints.openai import ( from app.api.endpoints.openai import (
MODEL_ID, MODEL_ID,
_CollectingMoviePilotAgent, _CollectingMoviePilotAgent,
_error_response as _openai_error_response,
) )
from app.api.openai_utils import build_anthropic_messages, build_prompt, build_session_id from app.api.openai_utils import (
build_anthropic_messages,
build_prompt,
build_session_id,
)
from app.core.config import settings from app.core.config import settings
from app.core.security import anthropic_api_key_header from app.core.security import anthropic_api_key_header
from app.schemas.types import MessageChannel from app.schemas.types import MessageChannel
@@ -91,7 +93,11 @@ async def _stream_anthropic_response(
pass pass
@router.post("/messages", summary="Anthropic compatible messages", response_model=schemas.AnthropicMessagesResponse) @router.post(
"/messages",
summary="Anthropic compatible messages",
response_model=schemas.AnthropicMessagesResponse,
)
async def messages( async def messages(
payload: schemas.AnthropicMessagesRequest, payload: schemas.AnthropicMessagesRequest,
x_api_key: Optional[str] = Security(anthropic_api_key_header), x_api_key: Optional[str] = Security(anthropic_api_key_header),
+36 -14
View File
@@ -10,11 +10,17 @@ from app.core.security import verify_token
router = APIRouter() router = APIRouter()
@router.get("/credits/{bangumiid}", summary="查询Bangumi演职员表", response_model=List[schemas.MediaPerson]) @router.get(
async def bangumi_credits(bangumiid: int, "/credits/{bangumiid}",
summary="查询Bangumi演职员表",
response_model=List[schemas.MediaPerson],
)
async def bangumi_credits(
bangumiid: int,
page: Optional[int] = 1, page: Optional[int] = 1,
count: Optional[int] = 20, count: Optional[int] = 20,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
查询Bangumi演职员表 查询Bangumi演职员表
""" """
@@ -24,11 +30,17 @@ async def bangumi_credits(bangumiid: int,
return [] return []
@router.get("/recommend/{bangumiid}", summary="查询Bangumi推荐", response_model=List[schemas.MediaInfo]) @router.get(
async def bangumi_recommend(bangumiid: int, "/recommend/{bangumiid}",
summary="查询Bangumi推荐",
response_model=List[schemas.MediaInfo],
)
async def bangumi_recommend(
bangumiid: int,
page: Optional[int] = 1, page: Optional[int] = 1,
count: Optional[int] = 20, count: Optional[int] = 20,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
查询Bangumi推荐 查询Bangumi推荐
""" """
@@ -38,20 +50,29 @@ async def bangumi_recommend(bangumiid: int,
return [] return []
@router.get("/person/{person_id}", summary="人物详情", response_model=schemas.MediaPerson) @router.get(
async def bangumi_person(person_id: int, "/person/{person_id}", summary="人物详情", response_model=schemas.MediaPerson
_: schemas.TokenPayload = Depends(verify_token)) -> Any: )
async def bangumi_person(
person_id: int, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
根据人物ID查询人物详情 根据人物ID查询人物详情
""" """
return await BangumiChain().async_person_detail(person_id=person_id) return await BangumiChain().async_person_detail(person_id=person_id)
@router.get("/person/credits/{person_id}", summary="人物参演作品", response_model=List[schemas.MediaInfo]) @router.get(
async def bangumi_person_credits(person_id: int, "/person/credits/{person_id}",
summary="人物参演作品",
response_model=List[schemas.MediaInfo],
)
async def bangumi_person_credits(
person_id: int,
page: Optional[int] = 1, page: Optional[int] = 1,
count: Optional[int] = 20, count: Optional[int] = 20,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
根据人物ID查询人物参演作品 根据人物ID查询人物参演作品
""" """
@@ -62,8 +83,9 @@ async def bangumi_person_credits(person_id: int,
@router.get("/{bangumiid}", summary="查询Bangumi详情", response_model=schemas.MediaInfo) @router.get("/{bangumiid}", summary="查询Bangumi详情", response_model=schemas.MediaInfo)
async def bangumi_info(bangumiid: int, async def bangumi_info(
_: schemas.TokenPayload = Depends(verify_token)) -> Any: bangumiid: int, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
查询Bangumi详情 查询Bangumi详情
""" """
+42 -17
View File
@@ -18,11 +18,15 @@ router = APIRouter()
@router.get("/statistic", summary="媒体数量统计", response_model=schemas.Statistic) @router.get("/statistic", summary="媒体数量统计", response_model=schemas.Statistic)
def statistic(name: Optional[str] = None, _: schemas.TokenPayload = Depends(verify_token)) -> Any: def statistic(
name: Optional[str] = None, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
查询媒体数量统计信息 查询媒体数量统计信息
""" """
media_statistics: Optional[List[schemas.Statistic]] = DashboardChain().media_statistic(name) media_statistics: Optional[List[schemas.Statistic]] = (
DashboardChain().media_statistic(name)
)
if media_statistics: if media_statistics:
# 汇总各媒体库统计信息 # 汇总各媒体库统计信息
ret_statistic = schemas.Statistic() ret_statistic = schemas.Statistic()
@@ -42,7 +46,9 @@ def statistic(name: Optional[str] = None, _: schemas.TokenPayload = Depends(veri
return schemas.Statistic() return schemas.Statistic()
@router.get("/statistic2", summary="媒体数量统计(API_TOKEN", response_model=schemas.Statistic) @router.get(
"/statistic2", summary="媒体数量统计(API_TOKEN", response_model=schemas.Statistic
)
def statistic2(_: Annotated[str, Depends(verify_apitoken)]) -> Any: def statistic2(_: Annotated[str, Depends(verify_apitoken)]) -> Any:
""" """
查询媒体数量统计信息 API_TOKEN认证(?token=xxx 查询媒体数量统计信息 API_TOKEN认证(?token=xxx
@@ -65,13 +71,12 @@ def storage(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
if _usage: if _usage:
total += _usage.total total += _usage.total
available += _usage.available available += _usage.available
return schemas.Storage( return schemas.Storage(total_storage=total, used_storage=total - available)
total_storage=total,
used_storage=total - available
@router.get(
"/storage2", summary="本地存储空间(API_TOKEN", response_model=schemas.Storage
) )
@router.get("/storage2", summary="本地存储空间(API_TOKEN", response_model=schemas.Storage)
def storage2(_: Annotated[str, Depends(verify_apitoken)]) -> Any: def storage2(_: Annotated[str, Depends(verify_apitoken)]) -> Any:
""" """
查询本地存储空间信息 API_TOKEN认证(?token=xxx 查询本地存储空间信息 API_TOKEN认证(?token=xxx
@@ -88,13 +93,17 @@ def processes(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
@router.get("/downloader", summary="下载器信息", response_model=schemas.DownloaderInfo) @router.get("/downloader", summary="下载器信息", response_model=schemas.DownloaderInfo)
def downloader(name: Optional[str] = None, _: schemas.TokenPayload = Depends(verify_token)) -> Any: def downloader(
name: Optional[str] = None, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
查询下载器信息 查询下载器信息
""" """
# 下载目录空间 # 下载目录空间
download_dirs = DirectoryHelper().get_local_download_dirs() download_dirs = DirectoryHelper().get_local_download_dirs()
_, free_space = SystemUtils.space_usage([Path(d.download_path) for d in download_dirs]) _, free_space = SystemUtils.space_usage(
[Path(d.download_path) for d in download_dirs]
)
# 下载器信息 # 下载器信息
downloader_info = schemas.DownloaderInfo() downloader_info = schemas.DownloaderInfo()
transfer_infos = DashboardChain().downloader_info(name) transfer_infos = DashboardChain().downloader_info(name)
@@ -108,7 +117,11 @@ def downloader(name: Optional[str] = None, _: schemas.TokenPayload = Depends(ver
return downloader_info return downloader_info
@router.get("/downloader2", summary="下载器信息(API_TOKEN", response_model=schemas.DownloaderInfo) @router.get(
"/downloader2",
summary="下载器信息(API_TOKEN",
response_model=schemas.DownloaderInfo,
)
def downloader2(_: Annotated[str, Depends(verify_apitoken)]) -> Any: def downloader2(_: Annotated[str, Depends(verify_apitoken)]) -> Any:
""" """
查询下载器信息 API_TOKEN认证(?token=xxx 查询下载器信息 API_TOKEN认证(?token=xxx
@@ -124,7 +137,11 @@ async def schedule(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
return Scheduler().list() return Scheduler().list()
@router.get("/schedule2", summary="后台服务(API_TOKEN", response_model=List[schemas.ScheduleInfo]) @router.get(
"/schedule2",
summary="后台服务(API_TOKEN",
response_model=List[schemas.ScheduleInfo],
)
async def schedule2(_: Annotated[str, Depends(verify_apitoken)]) -> Any: async def schedule2(_: Annotated[str, Depends(verify_apitoken)]) -> Any:
""" """
查询下载器信息 API_TOKEN认证(?token=xxx 查询下载器信息 API_TOKEN认证(?token=xxx
@@ -133,9 +150,11 @@ async def schedule2(_: Annotated[str, Depends(verify_apitoken)]) -> Any:
@router.get("/transfer", summary="文件整理统计", response_model=List[int]) @router.get("/transfer", summary="文件整理统计", response_model=List[int])
async def transfer(days: Optional[int] = 7, async def transfer(
days: Optional[int] = 7,
db: Session = Depends(get_db), db: Session = Depends(get_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
查询文件整理统计信息 查询文件整理统计信息
""" """
@@ -167,7 +186,11 @@ def memory(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
return SystemUtils.memory_usage() return SystemUtils.memory_usage()
@router.get("/memory2", summary="获取当前内存使用量和使用率(API_TOKEN)", response_model=List[int]) @router.get(
"/memory2",
summary="获取当前内存使用量和使用率(API_TOKEN)",
response_model=List[int],
)
def memory2(_: Annotated[str, Depends(verify_apitoken)]) -> Any: def memory2(_: Annotated[str, Depends(verify_apitoken)]) -> Any:
""" """
获取当前内存使用率 API_TOKEN认证(?token=xxx 获取当前内存使用率 API_TOKEN认证(?token=xxx
@@ -183,7 +206,9 @@ def network(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
return SystemUtils.network_usage() return SystemUtils.network_usage()
@router.get("/network2", summary="获取当前网络流量(API_TOKEN", response_model=List[int]) @router.get(
"/network2", summary="获取当前网络流量(API_TOKEN", response_model=List[int]
)
def network2(_: Annotated[str, Depends(verify_apitoken)]) -> Any: def network2(_: Annotated[str, Depends(verify_apitoken)]) -> Any:
""" """
获取当前网络流量 API_TOKEN认证(?token=xxx 获取当前网络流量 API_TOKEN认证(?token=xxx
+52 -25
View File
@@ -14,7 +14,11 @@ from app.schemas.types import ChainEventType, MediaType
router = APIRouter() router = APIRouter()
@router.get("/source", summary="获取探索数据源", response_model=List[schemas.DiscoverMediaSource]) @router.get(
"/source",
summary="获取探索数据源",
response_model=List[schemas.DiscoverMediaSource],
)
def source(_: schemas.TokenPayload = Depends(verify_token)) -> Any: def source(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
""" """
获取探索数据源 获取探索数据源
@@ -31,53 +35,69 @@ def source(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
@router.get("/bangumi", summary="探索Bangumi", response_model=List[schemas.MediaInfo]) @router.get("/bangumi", summary="探索Bangumi", response_model=List[schemas.MediaInfo])
async def bangumi(type: Optional[int] = 2, async def bangumi(
type: Optional[int] = 2,
cat: Optional[int] = None, cat: Optional[int] = None,
sort: Optional[str] = 'rank', sort: Optional[str] = "rank",
year: Optional[str] = None, year: Optional[str] = None,
page: Optional[int] = 1, page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
探索Bangumi 探索Bangumi
""" """
medias = await BangumiChain().async_discover(type=type, cat=cat, sort=sort, year=year, medias = await BangumiChain().async_discover(
limit=count, offset=(page - 1) * count) type=type, cat=cat, sort=sort, year=year, limit=count, offset=(page - 1) * count
)
if medias: if medias:
return [media.to_dict() for media in medias] return [media.to_dict() for media in medias]
return [] return []
@router.get("/douban_movies", summary="探索豆瓣电影", response_model=List[schemas.MediaInfo]) @router.get(
async def douban_movies(sort: Optional[str] = "R", "/douban_movies", summary="探索豆瓣电影", response_model=List[schemas.MediaInfo]
)
async def douban_movies(
sort: Optional[str] = "R",
tags: Optional[str] = "", tags: Optional[str] = "",
page: Optional[int] = 1, page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
浏览豆瓣电影信息 浏览豆瓣电影信息
""" """
movies = await DoubanChain().async_douban_discover(mtype=MediaType.MOVIE, movies = await DoubanChain().async_douban_discover(
sort=sort, tags=tags, page=page, count=count) mtype=MediaType.MOVIE, sort=sort, tags=tags, page=page, count=count
)
return [media.to_dict() for media in movies] if movies else [] return [media.to_dict() for media in movies] if movies else []
@router.get("/douban_tvs", summary="探索豆瓣剧集", response_model=List[schemas.MediaInfo]) @router.get(
async def douban_tvs(sort: Optional[str] = "R", "/douban_tvs", summary="探索豆瓣剧集", response_model=List[schemas.MediaInfo]
)
async def douban_tvs(
sort: Optional[str] = "R",
tags: Optional[str] = "", tags: Optional[str] = "",
page: Optional[int] = 1, page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
浏览豆瓣剧集信息 浏览豆瓣剧集信息
""" """
tvs = await DoubanChain().async_douban_discover(mtype=MediaType.TV, tvs = await DoubanChain().async_douban_discover(
sort=sort, tags=tags, page=page, count=count) mtype=MediaType.TV, sort=sort, tags=tags, page=page, count=count
)
return [media.to_dict() for media in tvs] if tvs else [] return [media.to_dict() for media in tvs] if tvs else []
@router.get("/tmdb_movies", summary="探索TMDB电影", response_model=List[schemas.MediaInfo]) @router.get(
async def tmdb_movies(sort_by: Optional[str] = "popularity.desc", "/tmdb_movies", summary="探索TMDB电影", response_model=List[schemas.MediaInfo]
)
async def tmdb_movies(
sort_by: Optional[str] = "popularity.desc",
with_genres: Optional[str] = "", with_genres: Optional[str] = "",
with_original_language: Optional[str] = "", with_original_language: Optional[str] = "",
with_keywords: Optional[str] = "", with_keywords: Optional[str] = "",
@@ -86,11 +106,13 @@ async def tmdb_movies(sort_by: Optional[str] = "popularity.desc",
vote_count: Optional[int] = 0, vote_count: Optional[int] = 0,
release_date: Optional[str] = "", release_date: Optional[str] = "",
page: Optional[int] = 1, page: Optional[int] = 1,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
浏览TMDB电影信息 浏览TMDB电影信息
""" """
movies = await TmdbChain().async_tmdb_discover(mtype=MediaType.MOVIE, movies = await TmdbChain().async_tmdb_discover(
mtype=MediaType.MOVIE,
sort_by=sort_by, sort_by=sort_by,
with_genres=with_genres, with_genres=with_genres,
with_original_language=with_original_language, with_original_language=with_original_language,
@@ -99,12 +121,14 @@ async def tmdb_movies(sort_by: Optional[str] = "popularity.desc",
vote_average=vote_average, vote_average=vote_average,
vote_count=vote_count, vote_count=vote_count,
release_date=release_date, release_date=release_date,
page=page) page=page,
)
return [movie.to_dict() for movie in movies] if movies else [] return [movie.to_dict() for movie in movies] if movies else []
@router.get("/tmdb_tvs", summary="探索TMDB剧集", response_model=List[schemas.MediaInfo]) @router.get("/tmdb_tvs", summary="探索TMDB剧集", response_model=List[schemas.MediaInfo])
async def tmdb_tvs(sort_by: Optional[str] = "popularity.desc", async def tmdb_tvs(
sort_by: Optional[str] = "popularity.desc",
with_genres: Optional[str] = "", with_genres: Optional[str] = "",
with_original_language: Optional[str] = "", with_original_language: Optional[str] = "",
with_keywords: Optional[str] = "", with_keywords: Optional[str] = "",
@@ -113,11 +137,13 @@ async def tmdb_tvs(sort_by: Optional[str] = "popularity.desc",
vote_count: Optional[int] = 0, vote_count: Optional[int] = 0,
release_date: Optional[str] = "", release_date: Optional[str] = "",
page: Optional[int] = 1, page: Optional[int] = 1,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
浏览TMDB剧集信息 浏览TMDB剧集信息
""" """
tvs = await TmdbChain().async_tmdb_discover(mtype=MediaType.TV, tvs = await TmdbChain().async_tmdb_discover(
mtype=MediaType.TV,
sort_by=sort_by, sort_by=sort_by,
with_genres=with_genres, with_genres=with_genres,
with_original_language=with_original_language, with_original_language=with_original_language,
@@ -126,5 +152,6 @@ async def tmdb_tvs(sort_by: Optional[str] = "popularity.desc",
vote_average=vote_average, vote_average=vote_average,
vote_count=vote_count, vote_count=vote_count,
release_date=release_date, release_date=release_date,
page=page) page=page,
)
return [tv.to_dict() for tv in tvs] if tvs else [] return [tv.to_dict() for tv in tvs] if tvs else []
+34 -16
View File
@@ -11,19 +11,28 @@ from app.schemas import MediaType
router = APIRouter() router = APIRouter()
@router.get("/person/{person_id}", summary="人物详情", response_model=schemas.MediaPerson) @router.get(
async def douban_person(person_id: int, "/person/{person_id}", summary="人物详情", response_model=schemas.MediaPerson
_: schemas.TokenPayload = Depends(verify_token)) -> Any: )
async def douban_person(
person_id: int, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
根据人物ID查询人物详情 根据人物ID查询人物详情
""" """
return await DoubanChain().async_person_detail(person_id=person_id) return await DoubanChain().async_person_detail(person_id=person_id)
@router.get("/person/credits/{person_id}", summary="人物参演作品", response_model=List[schemas.MediaInfo]) @router.get(
async def douban_person_credits(person_id: int, "/person/credits/{person_id}",
summary="人物参演作品",
response_model=List[schemas.MediaInfo],
)
async def douban_person_credits(
person_id: int,
page: Optional[int] = 1, page: Optional[int] = 1,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
根据人物ID查询人物参演作品 根据人物ID查询人物参演作品
""" """
@@ -33,10 +42,14 @@ async def douban_person_credits(person_id: int,
return [] return []
@router.get("/credits/{doubanid}/{type_name}", summary="豆瓣演员阵容", response_model=List[schemas.MediaPerson]) @router.get(
async def douban_credits(doubanid: str, "/credits/{doubanid}/{type_name}",
type_name: str, summary="豆瓣演员阵容",
_: schemas.TokenPayload = Depends(verify_token)) -> Any: response_model=List[schemas.MediaPerson],
)
async def douban_credits(
doubanid: str, type_name: str, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
根据豆瓣ID查询演员阵容,type_name: 电影/电视剧 根据豆瓣ID查询演员阵容,type_name: 电影/电视剧
""" """
@@ -48,10 +61,14 @@ async def douban_credits(doubanid: str,
return [] return []
@router.get("/recommend/{doubanid}/{type_name}", summary="豆瓣推荐电影/电视剧", response_model=List[schemas.MediaInfo]) @router.get(
async def douban_recommend(doubanid: str, "/recommend/{doubanid}/{type_name}",
type_name: str, summary="豆瓣推荐电影/电视剧",
_: schemas.TokenPayload = Depends(verify_token)) -> Any: response_model=List[schemas.MediaInfo],
)
async def douban_recommend(
doubanid: str, type_name: str, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
根据豆瓣ID查询推荐电影/电视剧,type_name: 电影/电视剧 根据豆瓣ID查询推荐电影/电视剧,type_name: 电影/电视剧
""" """
@@ -68,8 +85,9 @@ async def douban_recommend(doubanid: str,
@router.get("/{doubanid}", summary="查询豆瓣详情", response_model=schemas.MediaInfo) @router.get("/{doubanid}", summary="查询豆瓣详情", response_model=schemas.MediaInfo)
async def douban_info(doubanid: str, async def douban_info(
_: schemas.TokenPayload = Depends(verify_token)) -> Any: doubanid: str, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
根据豆瓣ID查询豆瓣媒体信息 根据豆瓣ID查询豆瓣媒体信息
""" """
+48 -29
View File
@@ -19,8 +19,8 @@ router = APIRouter()
@router.get("/", summary="正在下载", response_model=List[schemas.DownloadingTorrent]) @router.get("/", summary="正在下载", response_model=List[schemas.DownloadingTorrent])
def current( def current(
name: Optional[str] = None, name: Optional[str] = None, _: schemas.TokenPayload = Depends(verify_token)
_: schemas.TokenPayload = Depends(verify_token)) -> Any: ) -> Any:
""" """
查询正在下载的任务 查询正在下载的任务
""" """
@@ -33,7 +33,8 @@ def download(
torrent_in: schemas.TorrentInfo, torrent_in: schemas.TorrentInfo,
downloader: Annotated[str | None, Body()] = None, downloader: Annotated[str | None, Body()] = None,
save_path: Annotated[str | None, Body()] = None, save_path: Annotated[str | None, Body()] = None,
current_user: User = Depends(get_current_active_user)) -> Any: current_user: User = Depends(get_current_active_user),
) -> Any:
""" """
添加下载任务(含媒体信息) 添加下载任务(含媒体信息)
""" """
@@ -49,20 +50,22 @@ def download(
torrentinfo.site_downloader = downloader torrentinfo.site_downloader = downloader
# 上下文 # 上下文
context = Context( context = Context(
meta_info=metainfo, meta_info=metainfo, media_info=mediainfo, torrent_info=torrentinfo
media_info=mediainfo, )
torrent_info=torrentinfo did = DownloadChain().download_single(
context=context,
username=current_user.name,
save_path=save_path,
source="Manual",
) )
did = DownloadChain().download_single(context=context, username=current_user.name,
save_path=save_path, source="Manual")
if not did: if not did:
return schemas.Response(success=False, message="任务添加失败") return schemas.Response(success=False, message="任务添加失败")
return schemas.Response(success=True, data={ return schemas.Response(success=True, data={"download_id": did})
"download_id": did
})
@router.post("/add", summary="添加下载(不含媒体信息)", response_model=schemas.Response) @router.post(
"/add", summary="添加下载(不含媒体信息)", response_model=schemas.Response
)
def add( def add(
torrent_in: schemas.TorrentInfo, torrent_in: schemas.TorrentInfo,
tmdbid: Annotated[int | None, Body()] = None, tmdbid: Annotated[int | None, Body()] = None,
@@ -70,7 +73,8 @@ def add(
downloader: Annotated[str | None, Body()] = None, downloader: Annotated[str | None, Body()] = None,
# 保存路径, 支持<storage>:<path>, 如rclone:/MP, smb:/server/share/Movies等 # 保存路径, 支持<storage>:<path>, 如rclone:/MP, smb:/server/share/Movies等
save_path: Annotated[str | None, Body()] = None, save_path: Annotated[str | None, Body()] = None,
current_user: User = Depends(get_current_active_user)) -> Any: current_user: User = Depends(get_current_active_user),
) -> Any:
""" """
添加下载任务(不含媒体信息) 添加下载任务(不含媒体信息)
""" """
@@ -95,24 +99,27 @@ def add(
torrentinfo.from_dict(torrent_in.model_dump()) torrentinfo.from_dict(torrent_in.model_dump())
# 上下文 # 上下文
context = Context( context = Context(
meta_info=metainfo, meta_info=metainfo, media_info=mediainfo, torrent_info=torrentinfo
media_info=mediainfo,
torrent_info=torrentinfo
) )
did = DownloadChain().download_single(context=context, username=current_user.name, did = DownloadChain().download_single(
downloader=downloader, save_path=save_path, source="Manual") context=context,
username=current_user.name,
downloader=downloader,
save_path=save_path,
source="Manual",
)
if not did: if not did:
return schemas.Response(success=False, message="任务添加失败") return schemas.Response(success=False, message="任务添加失败")
return schemas.Response(success=True, data={ return schemas.Response(success=True, data={"download_id": did})
"download_id": did
})
@router.get("/start/{hashString}", summary="开始任务", response_model=schemas.Response) @router.get("/start/{hashString}", summary="开始任务", response_model=schemas.Response)
def start( def start(
hashString: str, name: Optional[str] = None, hashString: str,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: name: Optional[str] = None,
_: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
开如下载任务 开如下载任务
""" """
@@ -121,8 +128,11 @@ def start(
@router.get("/stop/{hashString}", summary="暂停任务", response_model=schemas.Response) @router.get("/stop/{hashString}", summary="暂停任务", response_model=schemas.Response)
def stop(hashString: str, name: Optional[str] = None, def stop(
_: schemas.TokenPayload = Depends(verify_token)) -> Any: hashString: str,
name: Optional[str] = None,
_: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
暂停下载任务 暂停下载任务
""" """
@@ -137,11 +147,17 @@ async def clients(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
""" """
downloaders: List[dict] = SystemConfigOper().get(SystemConfigKey.Downloaders) downloaders: List[dict] = SystemConfigOper().get(SystemConfigKey.Downloaders)
if downloaders: if downloaders:
return [{"name": d.get("name"), "type": d.get("type")} for d in downloaders if d.get("enabled")] return [
{"name": d.get("name"), "type": d.get("type")}
for d in downloaders
if d.get("enabled")
]
return [] return []
@router.get("/paths", summary="查询可用下载路径", response_model=List[schemas.DownloadDirectory]) @router.get(
"/paths", summary="查询可用下载路径", response_model=List[schemas.DownloadDirectory]
)
def paths(_: schemas.TokenPayload = Depends(verify_token)) -> Any: def paths(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
""" """
查询可直接用于下载接口 save_path 参数的下载路径 查询可直接用于下载接口 save_path 参数的下载路径
@@ -165,8 +181,11 @@ def paths(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
@router.delete("/{hashString}", summary="删除下载任务", response_model=schemas.Response) @router.delete("/{hashString}", summary="删除下载任务", response_model=schemas.Response)
def delete(hashString: str, name: Optional[str] = None, def delete(
_: schemas.TokenPayload = Depends(verify_token)) -> Any: hashString: str,
name: Optional[str] = None,
_: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
删除下载任务 删除下载任务
""" """
+59 -27
View File
@@ -18,7 +18,10 @@ from app.db import get_async_db, get_db
from app.db.models import User from app.db.models import User
from app.db.models.downloadhistory import DownloadHistory, DownloadFiles from app.db.models.downloadhistory import DownloadHistory, DownloadFiles
from app.db.models.transferhistory import TransferHistory from app.db.models.transferhistory import TransferHistory
from app.db.user_oper import get_current_active_superuser_async, get_current_active_superuser from app.db.user_oper import (
get_current_active_superuser_async,
get_current_active_superuser,
)
from app.helper.progress import ProgressHelper from app.helper.progress import ProgressHelper
from app.schemas.types import EventType from app.schemas.types import EventType
@@ -34,7 +37,9 @@ def normalize_history_ids(history_ids: list[int]) -> list[int]:
return normalized_ids return normalized_ids
def build_manual_redo_template_context(history: TransferHistory) -> dict[str, int | str]: def build_manual_redo_template_context(
history: TransferHistory,
) -> dict[str, int | str]:
"""仅负责把整理历史对象映射成 System Tasks 需要的模板变量。""" """仅负责把整理历史对象映射成 System Tasks 需要的模板变量。"""
src_fileitem = history.src_fileitem or {} src_fileitem = history.src_fileitem or {}
source_path = src_fileitem.get("path") if isinstance(src_fileitem, dict) else "" source_path = src_fileitem.get("path") if isinstance(src_fileitem, dict) else ""
@@ -200,11 +205,17 @@ def _start_batch_ai_redo_task(
asyncio.run_coroutine_threadsafe(runner(), global_vars.loop) asyncio.run_coroutine_threadsafe(runner(), global_vars.loop)
@router.get("/download", summary="查询下载历史记录", response_model=List[schemas.DownloadHistory]) @router.get(
async def download_history(page: Optional[int] = 1, "/download",
summary="查询下载历史记录",
response_model=List[schemas.DownloadHistory],
)
async def download_history(
page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
查询下载历史记录 查询下载历史记录
""" """
@@ -212,9 +223,11 @@ async def download_history(page: Optional[int] = 1,
@router.delete("/download", summary="删除下载历史记录", response_model=schemas.Response) @router.delete("/download", summary="删除下载历史记录", response_model=schemas.Response)
async def delete_download_history(history_in: schemas.DownloadHistory, async def delete_download_history(
history_in: schemas.DownloadHistory,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
删除下载历史记录 删除下载历史记录
""" """
@@ -223,12 +236,14 @@ async def delete_download_history(history_in: schemas.DownloadHistory,
@router.get("/transfer", summary="查询整理记录", response_model=schemas.Response) @router.get("/transfer", summary="查询整理记录", response_model=schemas.Response)
async def transfer_history(title: Optional[str] = None, async def transfer_history(
title: Optional[str] = None,
page: Optional[int] = 1, page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
status: Optional[bool] = None, status: Optional[bool] = None,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
查询整理记录 查询整理记录
""" """
@@ -242,26 +257,35 @@ async def transfer_history(title: Optional[str] = None,
if title: if title:
words = jieba.cut(title, HMM=False) words = jieba.cut(title, HMM=False)
title = "%".join(words) title = "%".join(words)
total = await TransferHistory.async_count_by_title(db, title=title, status=status) total = await TransferHistory.async_count_by_title(
result = await TransferHistory.async_list_by_title(db, title=title, page=page, db, title=title, status=status
count=count, status=status) )
result = await TransferHistory.async_list_by_title(
db, title=title, page=page, count=count, status=status
)
else: else:
result = await TransferHistory.async_list_by_page(db, page=page, count=count, status=status) result = await TransferHistory.async_list_by_page(
db, page=page, count=count, status=status
)
total = await TransferHistory.async_count(db, status=status) total = await TransferHistory.async_count(db, status=status)
return schemas.Response(success=True, return schemas.Response(
success=True,
data={ data={
"list": [item.to_dict() for item in result], "list": [item.to_dict() for item in result],
"total": total, "total": total,
}) },
)
@router.delete("/transfer", summary="删除整理记录", response_model=schemas.Response) @router.delete("/transfer", summary="删除整理记录", response_model=schemas.Response)
def delete_transfer_history(history_in: schemas.TransferHistory, def delete_transfer_history(
history_in: schemas.TransferHistory,
deletesrc: Optional[bool] = False, deletesrc: Optional[bool] = False,
deletedest: Optional[bool] = False, deletedest: Optional[bool] = False,
db: Session = Depends(get_db), db: Session = Depends(get_db),
_: User = Depends(get_current_active_superuser)) -> Any: _: User = Depends(get_current_active_superuser),
) -> Any:
""" """
删除整理记录 删除整理记录
""" """
@@ -278,23 +302,26 @@ def delete_transfer_history(history_in: schemas.TransferHistory,
src_fileitem = schemas.FileItem(**history.src_fileitem) src_fileitem = schemas.FileItem(**history.src_fileitem)
state = StorageChain().delete_media_file(src_fileitem) state = StorageChain().delete_media_file(src_fileitem)
if not state: if not state:
return schemas.Response(success=False, message=f"{src_fileitem.path} 删除失败") return schemas.Response(
success=False, message=f"{src_fileitem.path} 删除失败"
)
# 删除下载记录中关联的文件 # 删除下载记录中关联的文件
DownloadFiles.delete_by_fullpath(db, Path(src_fileitem.path).as_posix()) DownloadFiles.delete_by_fullpath(db, Path(src_fileitem.path).as_posix())
# 发送事件 # 发送事件
eventmanager.send_event( eventmanager.send_event(
EventType.DownloadFileDeleted, EventType.DownloadFileDeleted,
{ {"src": history.src, "hash": history.download_hash},
"src": history.src,
"hash": history.download_hash
}
) )
# 删除记录 # 删除记录
TransferHistory.delete(db, history_in.id) TransferHistory.delete(db, history_in.id)
return schemas.Response(success=True) return schemas.Response(success=True)
@router.post("/transfer/{history_id}/ai-redo", summary="智能助手重新整理", response_model=schemas.Response) @router.post(
"/transfer/{history_id}/ai-redo",
summary="智能助手重新整理",
response_model=schemas.Response,
)
def ai_redo_transfer_history( def ai_redo_transfer_history(
history_id: int, history_id: int,
db: Session = Depends(get_db), db: Session = Depends(get_db),
@@ -321,7 +348,9 @@ def ai_redo_transfer_history(
return schemas.Response(success=True, data={"progress_key": progress_key}) return schemas.Response(success=True, data={"progress_key": progress_key})
@router.post("/transfer/ai-redo", summary="智能助手批量重新整理", response_model=schemas.Response) @router.post(
"/transfer/ai-redo", summary="智能助手批量重新整理", response_model=schemas.Response
)
def batch_ai_redo_transfer_history( def batch_ai_redo_transfer_history(
payload: schemas.BatchTransferHistoryRedoRequest, payload: schemas.BatchTransferHistoryRedoRequest,
db: Session = Depends(get_db), db: Session = Depends(get_db),
@@ -349,7 +378,8 @@ def batch_ai_redo_transfer_history(
if missing_ids: if missing_ids:
return schemas.Response( return schemas.Response(
success=False, success=False,
message="整理记录不存在: " + ", ".join(str(history_id) for history_id in missing_ids), message="整理记录不存在: "
+ ", ".join(str(history_id) for history_id in missing_ids),
) )
prompt = build_batch_manual_redo_prompt(histories) prompt = build_batch_manual_redo_prompt(histories)
@@ -367,8 +397,10 @@ def batch_ai_redo_transfer_history(
@router.get("/empty/transfer", summary="清空整理记录", response_model=schemas.Response) @router.get("/empty/transfer", summary="清空整理记录", response_model=schemas.Response)
async def empty_transfer_history(db: AsyncSession = Depends(get_async_db), async def empty_transfer_history(
_: User = Depends(get_current_active_superuser_async)) -> Any: db: AsyncSession = Depends(get_async_db),
_: User = Depends(get_current_active_superuser_async),
) -> Any:
""" """
清空整理记录 清空整理记录
""" """
+5 -4
View File
@@ -125,7 +125,9 @@ async def start_llm_provider_auth(
callback_url = None callback_url = None
if payload.provider == "chatgpt" and payload.method == "browser_oauth": if payload.provider == "chatgpt" and payload.method == "browser_oauth":
callback_url = str( callback_url = str(
request.url_for("llm_provider_auth_callback", provider_id=payload.provider) request.url_for(
"llm_provider_auth_callback", provider_id=payload.provider
)
) )
result = await LLMProviderManager().start_auth( result = await LLMProviderManager().start_auth(
payload.provider, payload.provider,
@@ -250,9 +252,8 @@ async def llm_test(
if not payload.enabled: if not payload.enabled:
return schemas.Response(success=False, message="请先启用智能助手", data=data) return schemas.Response(success=False, message="请先启用智能助手", data=data)
if ( if payload.provider not in {"chatgpt", "github-copilot"} and (
payload.provider not in {"chatgpt", "github-copilot"} not payload.api_key or not payload.api_key.strip()
and (not payload.api_key or not payload.api_key.strip())
): ):
return schemas.Response( return schemas.Response(
success=False, success=False,
+12 -12
View File
@@ -19,14 +19,14 @@ router = APIRouter()
@router.post("/access-token", summary="获取token", response_model=schemas.Token) @router.post("/access-token", summary="获取token", response_model=schemas.Token)
def login_access_token( def login_access_token(
form_data: Annotated[OAuth2PasswordRequestForm, Depends()], form_data: Annotated[OAuth2PasswordRequestForm, Depends()],
otp_password: Annotated[str | None, Form()] = None otp_password: Annotated[str | None, Form()] = None,
) -> Any: ) -> Any:
""" """
获取认证Token 获取认证Token
""" """
success, user_or_message = UserChain().user_authenticate(username=form_data.username, success, user_or_message = UserChain().user_authenticate(
password=form_data.password, username=form_data.username, password=form_data.password, mfa_code=otp_password
mfa_code=otp_password) )
if not success: if not success:
# 如果是需要MFA验证,返回特殊标识 # 如果是需要MFA验证,返回特殊标识
@@ -34,21 +34,24 @@ def login_access_token(
raise HTTPException( raise HTTPException(
status_code=401, status_code=401,
detail="需要双重验证,请提供验证码或使用通行密钥", detail="需要双重验证,请提供验证码或使用通行密钥",
headers={"X-MFA-Required": "true"} headers={"X-MFA-Required": "true"},
) )
raise HTTPException(status_code=401, detail="用户名或密码错误") raise HTTPException(status_code=401, detail="用户名或密码错误")
# 用户等级 # 用户等级
level = SitesHelper().auth_level level = SitesHelper().auth_level
# 是否显示配置向导 # 是否显示配置向导
show_wizard = not SystemConfigOper().get(SystemConfigKey.SetupWizardState) and not settings.ADVANCED_MODE show_wizard = (
not SystemConfigOper().get(SystemConfigKey.SetupWizardState)
and not settings.ADVANCED_MODE
)
return schemas.Token( return schemas.Token(
access_token=security.create_access_token( access_token=security.create_access_token(
userid=user_or_message.id, userid=user_or_message.id,
username=user_or_message.name, username=user_or_message.name,
super_user=user_or_message.is_superuser, super_user=user_or_message.is_superuser,
expires_delta=timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES), expires_delta=timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES),
level=level level=level,
), ),
token_type="bearer", token_type="bearer",
super_user=user_or_message.is_superuser, super_user=user_or_message.is_superuser,
@@ -57,7 +60,7 @@ def login_access_token(
avatar=user_or_message.avatar, avatar=user_or_message.avatar,
level=level, level=level,
permissions=user_or_message.permissions or {}, permissions=user_or_message.permissions or {},
wizard=show_wizard wizard=show_wizard,
) )
@@ -68,10 +71,7 @@ def wallpaper() -> Any:
""" """
url = WallpaperHelper().get_wallpaper() url = WallpaperHelper().get_wallpaper()
if url: if url:
return schemas.Response( return schemas.Response(success=True, message=url)
success=True,
message=url
)
return schemas.Response(success=False) return schemas.Response(success=False)
+50 -66
View File
@@ -33,35 +33,32 @@ def list_exposed_tools():
获取 MCP 可见工具列表 获取 MCP 可见工具列表
""" """
return [ return [
tool for tool in moviepilot_tool_manager.list_tools() tool
for tool in moviepilot_tool_manager.list_tools()
if tool.name not in MCP_HIDDEN_TOOLS if tool.name not in MCP_HIDDEN_TOOLS
] ]
def create_jsonrpc_response(request_id: Union[str, int, None], result: Any) -> Dict[str, Any]: def create_jsonrpc_response(
request_id: Union[str, int, None], result: Any
) -> Dict[str, Any]:
""" """
创建 JSON-RPC 成功响应 创建 JSON-RPC 成功响应
""" """
response = { response = {"jsonrpc": "2.0", "id": request_id, "result": result}
"jsonrpc": "2.0",
"id": request_id,
"result": result
}
return response return response
def create_jsonrpc_error(request_id: Union[str, int, None], code: int, message: str, data: Any = None) -> Dict[ def create_jsonrpc_error(
str, Any]: request_id: Union[str, int, None], code: int, message: str, data: Any = None
) -> Dict[str, Any]:
""" """
创建 JSON-RPC 错误响应 创建 JSON-RPC 错误响应
""" """
error = { error = {
"jsonrpc": "2.0", "jsonrpc": "2.0",
"id": request_id, "id": request_id,
"error": { "error": {"code": code, "message": message},
"code": code,
"message": message
}
} }
if data is not None: if data is not None:
error["error"]["data"] = data error["error"]["data"] = data
@@ -70,8 +67,7 @@ def create_jsonrpc_error(request_id: Union[str, int, None], code: int, message:
@router.post("", summary="MCP JSON-RPC 端点", response_model=None) @router.post("", summary="MCP JSON-RPC 端点", response_model=None)
async def mcp_jsonrpc( async def mcp_jsonrpc(
request: Request, request: Request, _: Annotated[str, Depends(verify_apikey)] = None
_: Annotated[str, Depends(verify_apikey)] = None
) -> Union[JSONResponse, Response]: ) -> Union[JSONResponse, Response]:
""" """
MCP 标准 JSON-RPC 2.0 端点 MCP 标准 JSON-RPC 2.0 端点
@@ -84,14 +80,14 @@ async def mcp_jsonrpc(
logger.error(f"解析请求体失败: {e}") logger.error(f"解析请求体失败: {e}")
return JSONResponse( return JSONResponse(
status_code=400, status_code=400,
content=create_jsonrpc_error(None, -32700, "Parse error", str(e)) content=create_jsonrpc_error(None, -32700, "Parse error", str(e)),
) )
# 验证 JSON-RPC 格式 # 验证 JSON-RPC 格式
if not isinstance(body, dict) or body.get("jsonrpc") != "2.0": if not isinstance(body, dict) or body.get("jsonrpc") != "2.0":
return JSONResponse( return JSONResponse(
status_code=400, status_code=400,
content=create_jsonrpc_error(body.get("id"), -32600, "Invalid Request") content=create_jsonrpc_error(body.get("id"), -32600, "Invalid Request"),
) )
method = body.get("method") method = body.get("method")
@@ -114,7 +110,7 @@ async def mcp_jsonrpc(
else: else:
return JSONResponse( return JSONResponse(
status_code=400, status_code=400,
content={"error": "initialized must be a notification"} content={"error": "initialized must be a notification"},
) )
# 处理工具列表请求 # 处理工具列表请求
@@ -134,20 +130,22 @@ async def mcp_jsonrpc(
# 未知方法 # 未知方法
else: else:
return JSONResponse( return JSONResponse(
content=create_jsonrpc_error(request_id, -32601, f"Method not found: {method}") content=create_jsonrpc_error(
request_id, -32601, f"Method not found: {method}"
)
) )
except ValueError as e: except ValueError as e:
logger.warning(f"MCP 请求参数错误: {e}") logger.warning(f"MCP 请求参数错误: {e}")
return JSONResponse( return JSONResponse(
status_code=400, status_code=400,
content=create_jsonrpc_error(request_id, -32602, "Invalid params", str(e)) content=create_jsonrpc_error(request_id, -32602, "Invalid params", str(e)),
) )
except Exception as e: except Exception as e:
logger.error(f"处理 MCP 请求失败: {e}", exc_info=True) logger.error(f"处理 MCP 请求失败: {e}", exc_info=True)
return JSONResponse( return JSONResponse(
status_code=500, status_code=500,
content=create_jsonrpc_error(request_id, -32603, "Internal error", str(e)) content=create_jsonrpc_error(request_id, -32603, "Internal error", str(e)),
) )
@@ -158,7 +156,9 @@ async def handle_initialize(params: Dict[str, Any]) -> Dict[str, Any]:
protocol_version = params.get("protocolVersion") protocol_version = params.get("protocolVersion")
client_info = params.get("clientInfo", {}) client_info = params.get("clientInfo", {})
logger.info(f"MCP 初始化请求: 客户端={client_info.get('name')}, 协议版本={protocol_version}") logger.info(
f"MCP 初始化请求: 客户端={client_info.get('name')}, 协议版本={protocol_version}"
)
# 版本协商:选择客户端和服务器都支持的版本 # 版本协商:选择客户端和服务器都支持的版本
negotiated_version = MCP_PROTOCOL_VERSION negotiated_version = MCP_PROTOCOL_VERSION
@@ -168,7 +168,9 @@ async def handle_initialize(params: Dict[str, Any]) -> Dict[str, Any]:
logger.info(f"使用客户端协议版本: {negotiated_version}") logger.info(f"使用客户端协议版本: {negotiated_version}")
else: else:
# 客户端版本不支持,使用服务器默认版本 # 客户端版本不支持,使用服务器默认版本
logger.warning(f"协议版本不匹配: 客户端={protocol_version}, 使用服务器版本={negotiated_version}") logger.warning(
f"协议版本不匹配: 客户端={protocol_version}, 使用服务器版本={negotiated_version}"
)
return { return {
"protocolVersion": negotiated_version, "protocolVersion": negotiated_version,
@@ -176,14 +178,14 @@ async def handle_initialize(params: Dict[str, Any]) -> Dict[str, Any]:
"tools": { "tools": {
"listChanged": False # 暂不支持工具列表变更通知 "listChanged": False # 暂不支持工具列表变更通知
}, },
"logging": {} "logging": {},
}, },
"serverInfo": { "serverInfo": {
"name": "MoviePilot", "name": "MoviePilot",
"version": APP_VERSION, "version": APP_VERSION,
"description": "MoviePilot MCP Server - 电影自动化管理工具", "description": "MoviePilot MCP Server - 电影自动化管理工具",
}, },
"instructions": "MoviePilot MCP 服务器,提供媒体管理、订阅、下载等工具。" "instructions": "MoviePilot MCP 服务器,提供媒体管理、订阅、下载等工具。",
} }
@@ -199,13 +201,11 @@ async def handle_tools_list() -> Dict[str, Any]:
mcp_tool = { mcp_tool = {
"name": tool.name, "name": tool.name,
"description": tool.description, "description": tool.description,
"inputSchema": tool.input_schema "inputSchema": tool.input_schema,
} }
mcp_tools.append(mcp_tool) mcp_tools.append(mcp_tool)
return { return {"tools": mcp_tools}
"tools": mcp_tools
}
async def handle_tools_call(params: Dict[str, Any]) -> Dict[str, Any]: async def handle_tools_call(params: Dict[str, Any]) -> Dict[str, Any]:
@@ -224,30 +224,18 @@ async def handle_tools_call(params: Dict[str, Any]) -> Dict[str, Any]:
result_text = await moviepilot_tool_manager.call_tool(tool_name, arguments) result_text = await moviepilot_tool_manager.call_tool(tool_name, arguments)
return { return {"content": [{"type": "text", "text": result_text}]}
"content": [
{
"type": "text",
"text": result_text
}
]
}
except Exception as e: except Exception as e:
logger.error(f"工具调用失败: {tool_name}, 错误: {e}", exc_info=True) logger.error(f"工具调用失败: {tool_name}, 错误: {e}", exc_info=True)
return { return {
"content": [ "content": [{"type": "text", "text": f"错误: {str(e)}"}],
{ "isError": True,
"type": "text",
"text": f"错误: {str(e)}"
}
],
"isError": True
} }
@router.delete("", summary="终止 MCP 会话", response_model=None) @router.delete("", summary="终止 MCP 会话", response_model=None)
async def delete_mcp_session( async def delete_mcp_session(
_: Annotated[str, Depends(verify_apikey)] = None _: Annotated[str, Depends(verify_apikey)] = None,
) -> Union[JSONResponse, Response]: ) -> Union[JSONResponse, Response]:
""" """
终止 MCP 会话(无状态模式下仅返回成功) 终止 MCP 会话(无状态模式下仅返回成功)
@@ -257,10 +245,9 @@ async def delete_mcp_session(
# ==================== 兼容的 RESTful API 端点 ==================== # ==================== 兼容的 RESTful API 端点 ====================
@router.get("/tools", summary="列出所有可用工具", response_model=List[Dict[str, Any]]) @router.get("/tools", summary="列出所有可用工具", response_model=List[Dict[str, Any]])
async def list_tools( async def list_tools(_: Annotated[str, Depends(verify_apikey)]) -> Any:
_: Annotated[str, Depends(verify_apikey)]
) -> Any:
""" """
获取所有可用的工具列表 获取所有可用的工具列表
@@ -276,7 +263,7 @@ async def list_tools(
tool_dict = { tool_dict = {
"name": tool.name, "name": tool.name,
"description": tool.description, "description": tool.description,
"inputSchema": tool.input_schema "inputSchema": tool.input_schema,
} }
tools_list.append(tool_dict) tools_list.append(tool_dict)
@@ -288,8 +275,7 @@ async def list_tools(
@router.post("/tools/call", summary="调用工具", response_model=schemas.ToolCallResponse) @router.post("/tools/call", summary="调用工具", response_model=schemas.ToolCallResponse)
async def call_tool( async def call_tool(
request: schemas.ToolCallRequest, request: schemas.ToolCallRequest, _: Annotated[str, Depends(verify_apikey)] = None
_: Annotated[str, Depends(verify_apikey)] = None
) -> Any: ) -> Any:
""" """
调用指定的工具 调用指定的工具
@@ -301,24 +287,19 @@ async def call_tool(
if request.tool_name in MCP_HIDDEN_TOOLS: if request.tool_name in MCP_HIDDEN_TOOLS:
raise ValueError(f"工具 '{request.tool_name}' 未找到") raise ValueError(f"工具 '{request.tool_name}' 未找到")
result_text = await moviepilot_tool_manager.call_tool(request.tool_name, request.arguments) result_text = await moviepilot_tool_manager.call_tool(
request.tool_name, request.arguments
return schemas.ToolCallResponse(
success=True,
result=result_text
) )
return schemas.ToolCallResponse(success=True, result=result_text)
except Exception as e: except Exception as e:
logger.error(f"调用工具 {request.tool_name} 失败: {e}", exc_info=True) logger.error(f"调用工具 {request.tool_name} 失败: {e}", exc_info=True)
return schemas.ToolCallResponse( return schemas.ToolCallResponse(success=False, error=f"调用工具失败: {str(e)}")
success=False,
error=f"调用工具失败: {str(e)}"
)
@router.get("/tools/{tool_name}", summary="获取工具详情", response_model=Dict[str, Any]) @router.get("/tools/{tool_name}", summary="获取工具详情", response_model=Dict[str, Any])
async def get_tool_info( async def get_tool_info(
tool_name: str, tool_name: str, _: Annotated[str, Depends(verify_apikey)]
_: Annotated[str, Depends(verify_apikey)]
) -> Any: ) -> Any:
""" """
获取指定工具的详细信息 获取指定工具的详细信息
@@ -336,7 +317,7 @@ async def get_tool_info(
return { return {
"name": tool.name, "name": tool.name,
"description": tool.description, "description": tool.description,
"inputSchema": tool.input_schema "inputSchema": tool.input_schema,
} }
raise HTTPException(status_code=404, detail=f"工具 '{tool_name}' 未找到") raise HTTPException(status_code=404, detail=f"工具 '{tool_name}' 未找到")
@@ -347,10 +328,13 @@ async def get_tool_info(
raise HTTPException(status_code=500, detail=f"获取工具信息失败: {str(e)}") raise HTTPException(status_code=500, detail=f"获取工具信息失败: {str(e)}")
@router.get("/tools/{tool_name}/schema", summary="获取工具参数Schema", response_model=Dict[str, Any]) @router.get(
"/tools/{tool_name}/schema",
summary="获取工具参数Schema",
response_model=Dict[str, Any],
)
async def get_tool_schema( async def get_tool_schema(
tool_name: str, tool_name: str, _: Annotated[str, Depends(verify_apikey)]
_: Annotated[str, Depends(verify_apikey)]
) -> Any: ) -> Any:
""" """
获取指定工具的参数SchemaJSON Schema格式) 获取指定工具的参数SchemaJSON Schema格式)
+109 -44
View File
@@ -20,10 +20,14 @@ from app.schemas.types import ChainEventType
router = APIRouter() router = APIRouter()
@router.get("/recognize", summary="识别媒体信息(种子)", response_model=schemas.Context) @router.get(
async def recognize(title: str, "/recognize", summary="识别媒体信息(种子)", response_model=schemas.Context
)
async def recognize(
title: str,
subtitle: Optional[str] = None, subtitle: Optional[str] = None,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
根据标题、副标题识别媒体信息 根据标题、副标题识别媒体信息
""" """
@@ -35,10 +39,15 @@ async def recognize(title: str,
return schemas.Context() return schemas.Context()
@router.get("/recognize2", summary="识别种子媒体信息(API_TOKEN", response_model=schemas.Context) @router.get(
async def recognize2(_: Annotated[str, Depends(verify_apitoken)], "/recognize2",
summary="识别种子媒体信息(API_TOKEN",
response_model=schemas.Context,
)
async def recognize2(
_: Annotated[str, Depends(verify_apitoken)],
title: str, title: str,
subtitle: Optional[str] = None subtitle: Optional[str] = None,
) -> Any: ) -> Any:
""" """
根据标题、副标题识别媒体信息 API_TOKEN认证(?token=xxx 根据标题、副标题识别媒体信息 API_TOKEN认证(?token=xxx
@@ -47,9 +56,12 @@ async def recognize2(_: Annotated[str, Depends(verify_apitoken)],
return await recognize(title, subtitle) return await recognize(title, subtitle)
@router.get("/recognize_file", summary="识别媒体信息(文件)", response_model=schemas.Context) @router.get(
async def recognize_file(path: str, "/recognize_file", summary="识别媒体信息(文件)", response_model=schemas.Context
_: schemas.TokenPayload = Depends(verify_token)) -> Any: )
async def recognize_file(
path: str, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
根据文件路径识别媒体信息 根据文件路径识别媒体信息
""" """
@@ -60,9 +72,14 @@ async def recognize_file(path: str,
return schemas.Context() return schemas.Context()
@router.get("/recognize_file2", summary="识别文件媒体信息(API_TOKEN", response_model=schemas.Context) @router.get(
async def recognize_file2(path: str, "/recognize_file2",
_: Annotated[str, Depends(verify_apitoken)]) -> Any: summary="识别文件媒体信息(API_TOKEN",
response_model=schemas.Context,
)
async def recognize_file2(
path: str, _: Annotated[str, Depends(verify_apitoken)]
) -> Any:
""" """
根据文件路径识别媒体信息 API_TOKEN认证(?token=xxx 根据文件路径识别媒体信息 API_TOKEN认证(?token=xxx
""" """
@@ -71,11 +88,13 @@ async def recognize_file2(path: str,
@router.get("/search", summary="搜索媒体/人物信息", response_model=List[dict]) @router.get("/search", summary="搜索媒体/人物信息", response_model=List[dict])
async def search(title: str, async def search(
title: str,
type: Optional[str] = "media", type: Optional[str] = "media",
page: int = 1, page: int = 1,
count: int = 8, count: int = 8,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
模糊搜索媒体/人物信息列表 media:媒体信息,person:人物信息 模糊搜索媒体/人物信息列表 media:媒体信息,person:人物信息
""" """
@@ -94,7 +113,9 @@ async def search(title: str,
result = [media.to_dict() for media in medias] if medias else [] result = [media.to_dict() for media in medias] if medias else []
elif type == "collection": elif type == "collection":
collections = await media_chain.async_search_collections(name=title) collections = await media_chain.async_search_collections(name=title)
result = [collection.to_dict() for collection in collections] if collections else [] result = (
[collection.to_dict() for collection in collections] if collections else []
)
else: # person else: # person
persons = await media_chain.async_search_persons(name=title) persons = await media_chain.async_search_persons(name=title)
result = [person.model_dump() for person in persons] if persons else [] result = [person.model_dump() for person in persons] if persons else []
@@ -103,17 +124,21 @@ async def search(title: str,
return [] return []
# 排序和分页 # 排序和分页
setting_order = settings.SEARCH_SOURCE.split(',') if settings.SEARCH_SOURCE else [] setting_order = settings.SEARCH_SOURCE.split(",") if settings.SEARCH_SOURCE else []
sort_order = {source: index for index, source in enumerate(setting_order)} sort_order = {source: index for index, source in enumerate(setting_order)}
sorted_result = sorted(result, key=lambda x: sort_order.get(__get_source(x), 4)) sorted_result = sorted(result, key=lambda x: sort_order.get(__get_source(x), 4))
return sorted_result[(page - 1) * count : page * count] return sorted_result[(page - 1) * count : page * count]
@router.post("/scrape/{storage}", summary="刮削媒体信息", response_model=schemas.Response) @router.post(
def scrape(fileitem: schemas.FileItem, "/scrape/{storage}", summary="刮削媒体信息", response_model=schemas.Response
)
def scrape(
fileitem: schemas.FileItem,
storage: Optional[str] = "local", storage: Optional[str] = "local",
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
刮削媒体信息 刮削媒体信息
""" """
@@ -132,12 +157,14 @@ def scrape(fileitem: schemas.FileItem,
fileitem=fileitem, fileitem=fileitem,
meta=context.meta_info, meta=context.meta_info,
mediainfo=context.media_info, mediainfo=context.media_info,
overwrite=True overwrite=True,
) )
return schemas.Response(success=True, message=f"{fileitem.path} 刮削完成") return schemas.Response(success=True, message=f"{fileitem.path} 刮削完成")
@router.get("/category/config", summary="获取分类策略配置", response_model=schemas.Response) @router.get(
"/category/config", summary="获取分类策略配置", response_model=schemas.Response
)
def get_category_config(_: User = Depends(get_current_active_user)): def get_category_config(_: User = Depends(get_current_active_user)):
""" """
获取分类策略配置 获取分类策略配置
@@ -146,8 +173,12 @@ def get_category_config(_: User = Depends(get_current_active_user)):
return schemas.Response(success=True, data=config.model_dump()) return schemas.Response(success=True, data=config.model_dump())
@router.post("/category/config", summary="保存分类策略配置", response_model=schemas.Response) @router.post(
def save_category_config(config: CategoryConfig, _: User = Depends(get_current_active_superuser)): "/category/config", summary="保存分类策略配置", response_model=schemas.Response
)
def save_category_config(
config: CategoryConfig, _: User = Depends(get_current_active_superuser)
):
""" """
保存分类策略配置 保存分类策略配置
""" """
@@ -165,8 +196,14 @@ async def category(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
return MediaChain().media_category() or {} return MediaChain().media_category() or {}
@router.get("/group/seasons/{episode_group}", summary="查询剧集组季信息", response_model=List[schemas.MediaSeason]) @router.get(
async def group_seasons(episode_group: str, _: schemas.TokenPayload = Depends(verify_token)) -> Any: "/group/seasons/{episode_group}",
summary="查询剧集组季信息",
response_model=List[schemas.MediaSeason],
)
async def group_seasons(
episode_group: str, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
查询剧集组季信息(themoviedb) 查询剧集组季信息(themoviedb)
""" """
@@ -178,18 +215,24 @@ async def groups(tmdbid: int, _: schemas.TokenPayload = Depends(verify_token)) -
""" """
查询媒体剧集组列表(themoviedb) 查询媒体剧集组列表(themoviedb)
""" """
mediainfo = await MediaChain().async_recognize_media(tmdbid=tmdbid, mtype=MediaType.TV) mediainfo = await MediaChain().async_recognize_media(
tmdbid=tmdbid, mtype=MediaType.TV
)
if not mediainfo: if not mediainfo:
return [] return []
return mediainfo.episode_groups return mediainfo.episode_groups
@router.get("/seasons", summary="查询媒体季信息", response_model=List[schemas.MediaSeason]) @router.get(
async def seasons(mediaid: Optional[str] = None, "/seasons", summary="查询媒体季信息", response_model=List[schemas.MediaSeason]
)
async def seasons(
mediaid: Optional[str] = None,
title: Optional[str] = None, title: Optional[str] = None,
year: str = None, year: str = None,
season: int = None, season: int = None,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
查询媒体季信息 查询媒体季信息
""" """
@@ -212,28 +255,39 @@ async def seasons(mediaid: Optional[str] = None,
) )
if mediainfo: if mediainfo:
if settings.RECOGNIZE_SOURCE == "themoviedb": if settings.RECOGNIZE_SOURCE == "themoviedb":
seasons_info = await TmdbChain().async_tmdb_seasons(tmdbid=mediainfo.tmdb_id) seasons_info = await TmdbChain().async_tmdb_seasons(
tmdbid=mediainfo.tmdb_id
)
if seasons_info: if seasons_info:
if season is not None: if season is not None:
return [sea for sea in seasons_info if sea.season_number == season] return [
sea for sea in seasons_info if sea.season_number == season
]
return seasons_info return seasons_info
else: else:
sea = season if season is not None else 1 sea = season if season is not None else 1
return [schemas.MediaSeason( return [
schemas.MediaSeason(
season_number=sea, season_number=sea,
poster_path=mediainfo.poster_path, poster_path=mediainfo.poster_path,
name=f"{sea}", name=f"{sea}",
air_date=mediainfo.release_date, air_date=mediainfo.release_date,
overview=mediainfo.overview, overview=mediainfo.overview,
vote_average=mediainfo.vote_average, vote_average=mediainfo.vote_average,
episode_count=mediainfo.number_of_episodes episode_count=mediainfo.number_of_episodes,
)] )
]
return [] return []
@router.get("/{mediaid}", summary="查询媒体详情", response_model=schemas.MediaInfo) @router.get("/{mediaid}", summary="查询媒体详情", response_model=schemas.MediaInfo)
async def detail(mediaid: str, type_name: str, title: Optional[str] = None, year: str = None, async def detail(
_: schemas.TokenPayload = Depends(verify_token)) -> Any: mediaid: str,
type_name: str,
title: Optional[str] = None,
year: str = None,
_: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
根据媒体ID查询themoviedb或豆瓣媒体信息,type_name: 电影/电视剧 根据媒体ID查询themoviedb或豆瓣媒体信息,type_name: 电影/电视剧
""" """
@@ -241,26 +295,37 @@ async def detail(mediaid: str, type_name: str, title: Optional[str] = None, year
mediainfo = None mediainfo = None
mediachain = MediaChain() mediachain = MediaChain()
if mediaid.startswith("tmdb:"): if mediaid.startswith("tmdb:"):
mediainfo = await mediachain.async_recognize_media(tmdbid=int(mediaid[5:]), mtype=mtype) mediainfo = await mediachain.async_recognize_media(
tmdbid=int(mediaid[5:]), mtype=mtype
)
elif mediaid.startswith("douban:"): elif mediaid.startswith("douban:"):
mediainfo = await mediachain.async_recognize_media(doubanid=mediaid[7:], mtype=mtype) mediainfo = await mediachain.async_recognize_media(
doubanid=mediaid[7:], mtype=mtype
)
elif mediaid.startswith("bangumi:"): elif mediaid.startswith("bangumi:"):
mediainfo = await mediachain.async_recognize_media(bangumiid=int(mediaid[8:]), mtype=mtype) mediainfo = await mediachain.async_recognize_media(
bangumiid=int(mediaid[8:]), mtype=mtype
)
else: else:
# 广播事件解析媒体信息 # 广播事件解析媒体信息
event_data = MediaRecognizeConvertEventData( event_data = MediaRecognizeConvertEventData(
mediaid=mediaid, mediaid=mediaid, convert_type=settings.RECOGNIZE_SOURCE
convert_type=settings.RECOGNIZE_SOURCE )
event = await eventmanager.async_send_event(
ChainEventType.MediaRecognizeConvert, event_data
) )
event = await eventmanager.async_send_event(ChainEventType.MediaRecognizeConvert, event_data)
# 使用事件返回的上下文数据 # 使用事件返回的上下文数据
if event and event.event_data and event.event_data.media_dict: if event and event.event_data and event.event_data.media_dict:
event_data: MediaRecognizeConvertEventData = event.event_data event_data: MediaRecognizeConvertEventData = event.event_data
new_id = event_data.media_dict.get("id") new_id = event_data.media_dict.get("id")
if event_data.convert_type == "themoviedb": if event_data.convert_type == "themoviedb":
mediainfo = await mediachain.async_recognize_media(tmdbid=new_id, mtype=mtype) mediainfo = await mediachain.async_recognize_media(
tmdbid=new_id, mtype=mtype
)
elif event_data.convert_type == "douban": elif event_data.convert_type == "douban":
mediainfo = await mediachain.async_recognize_media(doubanid=new_id, mtype=mtype) mediainfo = await mediachain.async_recognize_media(
doubanid=new_id, mtype=mtype
)
elif title: elif title:
# 使用名称识别兜底 # 使用名称识别兜底
meta = MetaInfo(title) meta = MetaInfo(title)
+63 -26
View File
@@ -30,8 +30,11 @@ def start_message_chain(body: Any, form: Any, args: Any):
@router.post("/", summary="接收用户消息", response_model=schemas.Response) @router.post("/", summary="接收用户消息", response_model=schemas.Response)
async def user_message(background_tasks: BackgroundTasks, request: Request, async def user_message(
_: schemas.TokenPayload = Depends(verify_apitoken)): background_tasks: BackgroundTasks,
request: Request,
_: schemas.TokenPayload = Depends(verify_apitoken),
):
""" """
用户消息响应,配置请求中需要添加参数:token=API_TOKEN&source=消息配置名 用户消息响应,配置请求中需要添加参数:token=API_TOKEN&source=消息配置名
""" """
@@ -106,10 +109,12 @@ async def web_message(
@router.get("/web", summary="获取WEB消息", response_model=List[dict]) @router.get("/web", summary="获取WEB消息", response_model=List[dict])
async def get_web_message(_: schemas.TokenPayload = Depends(verify_token), async def get_web_message(
_: schemas.TokenPayload = Depends(verify_token),
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
page: Optional[int] = 1, page: Optional[int] = 1,
count: Optional[int] = 20): count: Optional[int] = 20,
):
""" """
获取WEB消息列表 获取WEB消息列表
""" """
@@ -124,8 +129,13 @@ async def get_web_message(_: schemas.TokenPayload = Depends(verify_token),
return ret_messages return ret_messages
def wechat_verify(echostr: str, msg_signature: str, timestamp: Union[str, int], nonce: str, def wechat_verify(
source: Optional[str] = None) -> Any: echostr: str,
msg_signature: str,
timestamp: Union[str, int],
nonce: str,
source: Optional[str] = None,
) -> Any:
""" """
微信验证响应 微信验证响应
""" """
@@ -133,21 +143,31 @@ def wechat_verify(echostr: str, msg_signature: str, timestamp: Union[str, int],
client_configs = ServiceConfigHelper.get_notification_configs() client_configs = ServiceConfigHelper.get_notification_configs()
if not client_configs: if not client_configs:
return "未找到对应的消息配置" return "未找到对应的消息配置"
client_config = next((config for config in client_configs if client_config = next(
config.type == "wechat" (
config
for config in client_configs
if config.type == "wechat"
and config.enabled and config.enabled
and config.config.get("WECHAT_MODE", "app") != "bot" and config.config.get("WECHAT_MODE", "app") != "bot"
and (not source or config.name == source)), None) and (not source or config.name == source)
),
None,
)
if not client_config: if not client_config:
return "未找到对应的消息配置" return "未找到对应的消息配置"
try: try:
wxcpt = WXBizMsgCrypt(sToken=client_config.config.get('WECHAT_TOKEN'), wxcpt = WXBizMsgCrypt(
sEncodingAESKey=client_config.config.get('WECHAT_ENCODING_AESKEY'), sToken=client_config.config.get("WECHAT_TOKEN"),
sReceiveId=client_config.config.get('WECHAT_CORPID')) sEncodingAESKey=client_config.config.get("WECHAT_ENCODING_AESKEY"),
ret, sEchoStr = wxcpt.VerifyURL(sMsgSignature=msg_signature, sReceiveId=client_config.config.get("WECHAT_CORPID"),
)
ret, sEchoStr = wxcpt.VerifyURL(
sMsgSignature=msg_signature,
sTimeStamp=timestamp, sTimeStamp=timestamp,
sNonce=nonce, sNonce=nonce,
sEchoStr=echostr) sEchoStr=echostr,
)
if ret == 0: if ret == 0:
# 验证URL成功,将sEchoStr返回给企业号 # 验证URL成功,将sEchoStr返回给企业号
return PlainTextResponse(sEchoStr) return PlainTextResponse(sEchoStr)
@@ -165,21 +185,35 @@ def vocechat_verify() -> Any:
@router.get("/", summary="回调请求验证") @router.get("/", summary="回调请求验证")
def incoming_verify(token: Optional[str] = None, echostr: Optional[str] = None, msg_signature: Optional[str] = None, def incoming_verify(
timestamp: Union[str, int] = None, nonce: Optional[str] = None, source: Optional[str] = None, token: Optional[str] = None,
_: schemas.TokenPayload = Depends(verify_apitoken)) -> Any: echostr: Optional[str] = None,
msg_signature: Optional[str] = None,
timestamp: Union[str, int] = None,
nonce: Optional[str] = None,
source: Optional[str] = None,
_: schemas.TokenPayload = Depends(verify_apitoken),
) -> Any:
""" """
微信/VoceChat等验证响应 微信/VoceChat等验证响应
""" """
logger.info(f"收到验证请求: token={token}, echostr={echostr}, " logger.info(
f"msg_signature={msg_signature}, timestamp={timestamp}, nonce={nonce}") f"收到验证请求: token={token}, echostr={echostr}, "
f"msg_signature={msg_signature}, timestamp={timestamp}, nonce={nonce}"
)
if echostr and msg_signature and timestamp and nonce: if echostr and msg_signature and timestamp and nonce:
return wechat_verify(echostr, msg_signature, timestamp, nonce, source) return wechat_verify(echostr, msg_signature, timestamp, nonce, source)
return vocechat_verify() return vocechat_verify()
@router.post("/webpush/subscribe", summary="客户端webpush通知订阅", response_model=schemas.Response) @router.post(
async def subscribe(subscription: schemas.Subscription, _: schemas.TokenPayload = Depends(verify_token)): "/webpush/subscribe",
summary="客户端webpush通知订阅",
response_model=schemas.Response,
)
async def subscribe(
subscription: schemas.Subscription, _: schemas.TokenPayload = Depends(verify_token)
):
""" """
客户端webpush通知订阅 客户端webpush通知订阅
""" """
@@ -190,8 +224,13 @@ async def subscribe(subscription: schemas.Subscription, _: schemas.TokenPayload
return schemas.Response(success=True) return schemas.Response(success=True)
@router.post("/webpush/send", summary="发送webpush通知", response_model=schemas.Response) @router.post(
def send_notification(payload: schemas.SubscriptionMessage, _: schemas.TokenPayload = Depends(verify_token)): "/webpush/send", summary="发送webpush通知", response_model=schemas.Response
)
def send_notification(
payload: schemas.SubscriptionMessage,
_: schemas.TokenPayload = Depends(verify_token),
):
""" """
发送webpush通知 发送webpush通知
""" """
@@ -201,9 +240,7 @@ def send_notification(payload: schemas.SubscriptionMessage, _: schemas.TokenPayl
subscription_info=sub, subscription_info=sub,
data=json.dumps(payload.model_dump()), data=json.dumps(payload.model_dump()),
vapid_private_key=settings.VAPID.get("privateKey"), vapid_private_key=settings.VAPID.get("privateKey"),
vapid_claims={ vapid_claims={"sub": settings.VAPID.get("subject")},
"sub": settings.VAPID.get("subject")
},
) )
except WebPushException as err: except WebPushException as err:
logger.error(f"WebPush发送失败: {str(err)}") logger.error(f"WebPush发送失败: {str(err)}")
+156 -131
View File
@@ -2,6 +2,7 @@
MFA (Multi-Factor Authentication) API 端点 MFA (Multi-Factor Authentication) API 端点
包含 OTP 和 PassKey 相关功能 包含 OTP 和 PassKey 相关功能
""" """
from datetime import timedelta from datetime import timedelta
from typing import Any, Annotated, Optional from typing import Any, Annotated, Optional
@@ -26,6 +27,7 @@ router = APIRouter()
# ==================== 辅助函数 ==================== # ==================== 辅助函数 ====================
def _build_credential_list(passkeys: list[PassKey]) -> list[dict[str, Any]]: def _build_credential_list(passkeys: list[PassKey]) -> list[dict[str, Any]]:
""" """
构建凭证列表 构建凭证列表
@@ -33,13 +35,14 @@ def _build_credential_list(passkeys: list[PassKey]) -> list[dict[str, Any]]:
:param passkeys: PassKey 列表 :param passkeys: PassKey 列表
:return: 凭证字典列表 :return: 凭证字典列表
""" """
return [ return (
{ [
'credential_id': pk.credential_id, {"credential_id": pk.credential_id, "transports": pk.transports}
'transports': pk.transports
}
for pk in passkeys for pk in passkeys
] if passkeys else [] ]
if passkeys
else []
)
def _extract_and_standardize_credential_id(credential: dict) -> str: def _extract_and_standardize_credential_id(credential: dict) -> str:
@@ -50,16 +53,14 @@ def _extract_and_standardize_credential_id(credential: dict) -> str:
:return: 标准化后的 credential_id :return: 标准化后的 credential_id
:raises ValueError: 如果凭证无效 :raises ValueError: 如果凭证无效
""" """
credential_id_raw = credential.get('id') or credential.get('rawId') credential_id_raw = credential.get("id") or credential.get("rawId")
if not credential_id_raw: if not credential_id_raw:
raise ValueError("无效的凭证") raise ValueError("无效的凭证")
return PassKeyHelper.standardize_credential_id(credential_id_raw) return PassKeyHelper.standardize_credential_id(credential_id_raw)
def _verify_passkey_and_update( def _verify_passkey_and_update(
credential: dict, credential: dict, challenge: str, passkey: PassKey
challenge: str,
passkey: PassKey
) -> tuple[bool, int]: ) -> tuple[bool, int]:
""" """
验证 PassKey 并更新使用时间和签名计数 验证 PassKey 并更新使用时间和签名计数
@@ -73,7 +74,7 @@ def _verify_passkey_and_update(
credential=credential, credential=credential,
expected_challenge=challenge, expected_challenge=challenge,
credential_public_key=passkey.public_key, credential_public_key=passkey.public_key,
credential_current_sign_count=passkey.sign_count credential_current_sign_count=passkey.sign_count,
) )
if success: if success:
@@ -95,23 +96,35 @@ async def _check_user_has_passkey(db: AsyncSession, user_id: int) -> bool:
# ==================== 请求模型 ==================== # ==================== 请求模型 ====================
class OtpVerifyRequest(schemas.BaseModel): class OtpVerifyRequest(schemas.BaseModel):
"""OTP验证请求""" """OTP验证请求"""
uri: str uri: str
otpPassword: str otpPassword: str
class OtpDisableRequest(schemas.BaseModel): class OtpDisableRequest(schemas.BaseModel):
"""OTP禁用请求""" """OTP禁用请求"""
password: str password: str
class PassKeyDeleteRequest(schemas.BaseModel): class PassKeyDeleteRequest(schemas.BaseModel):
"""PassKey删除请求""" """PassKey删除请求"""
passkey_id: int passkey_id: int
password: str password: str
# ==================== 通用 MFA 接口 ==================== # ==================== 通用 MFA 接口 ====================
@router.get('/status/{username}', summary='判断用户是否开启双重验证(MFA)', response_model=schemas.Response)
@router.get(
"/status/{username}",
summary="判断用户是否开启双重验证(MFA)",
response_model=schemas.Response,
)
async def mfa_status(username: str, db: AsyncSession = Depends(get_async_db)) -> Any: async def mfa_status(username: str, db: AsyncSession = Depends(get_async_db)) -> Any:
""" """
检查指定用户是否启用了任何双重验证方式(OTP 或 PassKey 检查指定用户是否启用了任何双重验证方式(OTP 或 PassKey
@@ -132,33 +145,40 @@ async def mfa_status(username: str, db: AsyncSession = Depends(get_async_db)) ->
# ==================== OTP 相关接口 ==================== # ==================== OTP 相关接口 ====================
@router.post('/otp/generate', summary='生成 OTP 验证 URI', response_model=schemas.Response)
@router.post(
"/otp/generate", summary="生成 OTP 验证 URI", response_model=schemas.Response
)
def otp_generate( def otp_generate(
current_user: Annotated[User, Depends(get_current_active_user)] current_user: Annotated[User, Depends(get_current_active_user)],
) -> Any: ) -> Any:
"""生成 OTP 密钥及对应的 URI""" """生成 OTP 密钥及对应的 URI"""
secret, uri = OtpUtils.generate_secret_key(current_user.name) secret, uri = OtpUtils.generate_secret_key(current_user.name)
return schemas.Response(success=secret != "", data={'secret': secret, 'uri': uri}) return schemas.Response(success=secret != "", data={"secret": secret, "uri": uri})
@router.post('/otp/verify', summary='绑定并验证 OTP', response_model=schemas.Response) @router.post("/otp/verify", summary="绑定并验证 OTP", response_model=schemas.Response)
async def otp_verify( async def otp_verify(
data: OtpVerifyRequest, data: OtpVerifyRequest,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
current_user: User = Depends(get_current_active_user_async) current_user: User = Depends(get_current_active_user_async),
) -> Any: ) -> Any:
"""验证用户输入的 OTP 码,验证通过后正式开启 OTP 验证""" """验证用户输入的 OTP 码,验证通过后正式开启 OTP 验证"""
if not OtpUtils.is_legal(data.uri, data.otpPassword): if not OtpUtils.is_legal(data.uri, data.otpPassword):
return schemas.Response(success=False, message="验证码错误") return schemas.Response(success=False, message="验证码错误")
await current_user.async_update_otp_by_name(db, current_user.name, True, OtpUtils.get_secret(data.uri)) await current_user.async_update_otp_by_name(
db, current_user.name, True, OtpUtils.get_secret(data.uri)
)
return schemas.Response(success=True) return schemas.Response(success=True)
@router.post('/otp/disable', summary='关闭当前用户的 OTP 验证', response_model=schemas.Response) @router.post(
"/otp/disable", summary="关闭当前用户的 OTP 验证", response_model=schemas.Response
)
async def otp_disable( async def otp_disable(
data: OtpDisableRequest, data: OtpDisableRequest,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
current_user: User = Depends(get_current_active_user_async) current_user: User = Depends(get_current_active_user_async),
) -> Any: ) -> Any:
"""关闭当前用户的 OTP 验证功能""" """关闭当前用户的 OTP 验证功能"""
# 安全检查:如果存在 PassKey,默认不允许关闭 OTP,除非配置允许 # 安全检查:如果存在 PassKey,默认不允许关闭 OTP,除非配置允许
@@ -166,7 +186,7 @@ async def otp_disable(
if has_passkey and not settings.PASSKEY_ALLOW_REGISTER_WITHOUT_OTP: if has_passkey and not settings.PASSKEY_ALLOW_REGISTER_WITHOUT_OTP:
return schemas.Response( return schemas.Response(
success=False, success=False,
message="您已注册通行密钥,为了防止域名配置变更导致无法登录,请先删除所有通行密钥再关闭 OTP 验证" message="您已注册通行密钥,为了防止域名配置变更导致无法登录,请先删除所有通行密钥再关闭 OTP 验证",
) )
# 验证密码 # 验证密码
@@ -178,13 +198,16 @@ async def otp_disable(
# ==================== PassKey 相关接口 ==================== # ==================== PassKey 相关接口 ====================
class PassKeyRegistrationStart(schemas.BaseModel): class PassKeyRegistrationStart(schemas.BaseModel):
"""PassKey注册开始请求""" """PassKey注册开始请求"""
name: str = "通行密钥" name: str = "通行密钥"
class PassKeyRegistrationFinish(schemas.BaseModel): class PassKeyRegistrationFinish(schemas.BaseModel):
"""PassKey注册完成请求""" """PassKey注册完成请求"""
credential: dict credential: dict
challenge: str challenge: str
name: str = "通行密钥" name: str = "通行密钥"
@@ -192,18 +215,24 @@ class PassKeyRegistrationFinish(schemas.BaseModel):
class PassKeyAuthenticationStart(schemas.BaseModel): class PassKeyAuthenticationStart(schemas.BaseModel):
"""PassKey认证开始请求""" """PassKey认证开始请求"""
username: Optional[str] = None username: Optional[str] = None
class PassKeyAuthenticationFinish(schemas.BaseModel): class PassKeyAuthenticationFinish(schemas.BaseModel):
"""PassKey认证完成请求""" """PassKey认证完成请求"""
credential: dict credential: dict
challenge: str challenge: str
@router.post("/passkey/register/start", summary="开始注册 PassKey", response_model=schemas.Response) @router.post(
"/passkey/register/start",
summary="开始注册 PassKey",
response_model=schemas.Response,
)
def passkey_register_start( def passkey_register_start(
current_user: Annotated[User, Depends(get_current_active_user)] current_user: Annotated[User, Depends(get_current_active_user)],
) -> Any: ) -> Any:
"""开始注册 PassKey - 生成注册选项""" """开始注册 PassKey - 生成注册选项"""
try: try:
@@ -211,53 +240,59 @@ def passkey_register_start(
if not current_user.is_otp and not settings.PASSKEY_ALLOW_REGISTER_WITHOUT_OTP: if not current_user.is_otp and not settings.PASSKEY_ALLOW_REGISTER_WITHOUT_OTP:
return schemas.Response( return schemas.Response(
success=False, success=False,
message="为了确保在域名配置错误时仍能找回访问权限,请先启用 OTP 验证码再注册通行密钥" message="为了确保在域名配置错误时仍能找回访问权限,请先启用 OTP 验证码再注册通行密钥",
) )
# 获取用户已有的PassKey # 获取用户已有的PassKey
existing_passkeys = PassKey.get_by_user_id(db=None, user_id=current_user.id) existing_passkeys = PassKey.get_by_user_id(db=None, user_id=current_user.id)
existing_credentials = _build_credential_list(existing_passkeys) if existing_passkeys else None existing_credentials = (
_build_credential_list(existing_passkeys) if existing_passkeys else None
)
# 生成注册选项 # 生成注册选项
options_json, challenge = PassKeyHelper.generate_registration_options( options_json, challenge = PassKeyHelper.generate_registration_options(
user_id=current_user.id, user_id=current_user.id,
username=current_user.name, username=current_user.name,
display_name=current_user.settings.get('nickname') if current_user.settings else None, display_name=current_user.settings.get("nickname")
existing_credentials=existing_credentials if current_user.settings
else None,
existing_credentials=existing_credentials,
) )
return schemas.Response( return schemas.Response(
success=True, success=True, data={"options": options_json, "challenge": challenge}
data={
'options': options_json,
'challenge': challenge
}
) )
except Exception as e: except Exception as e:
logger.error(f"生成PassKey注册选项失败: {e}") logger.error(f"生成PassKey注册选项失败: {e}")
return schemas.Response( return schemas.Response(success=False, message=f"生成注册选项失败: {str(e)}")
success=False,
message=f"生成注册选项失败: {str(e)}"
@router.post(
"/passkey/register/finish",
summary="完成注册 PassKey",
response_model=schemas.Response,
) )
@router.post("/passkey/register/finish", summary="完成注册 PassKey", response_model=schemas.Response)
def passkey_register_finish( def passkey_register_finish(
passkey_req: PassKeyRegistrationFinish, passkey_req: PassKeyRegistrationFinish,
current_user: Annotated[User, Depends(get_current_active_user)] current_user: Annotated[User, Depends(get_current_active_user)],
) -> Any: ) -> Any:
"""完成注册 PassKey - 验证并保存凭证""" """完成注册 PassKey - 验证并保存凭证"""
try: try:
# 验证注册响应 # 验证注册响应
credential_id, public_key, sign_count, aaguid = PassKeyHelper.verify_registration_response( credential_id, public_key, sign_count, aaguid = (
PassKeyHelper.verify_registration_response(
credential=passkey_req.credential, credential=passkey_req.credential,
expected_challenge=passkey_req.challenge expected_challenge=passkey_req.challenge,
)
) )
# 提取transports # 提取transports
transports = None transports = None
if 'response' in passkey_req.credential and 'transports' in passkey_req.credential['response']: if (
transports = ','.join(passkey_req.credential['response']['transports']) "response" in passkey_req.credential
and "transports" in passkey_req.credential["response"]
):
transports = ",".join(passkey_req.credential["response"]["transports"])
# 保存到数据库 # 保存到数据库
passkey = PassKey( passkey = PassKey(
@@ -267,27 +302,25 @@ def passkey_register_finish(
sign_count=sign_count, sign_count=sign_count,
name=passkey_req.name or "通行密钥", name=passkey_req.name or "通行密钥",
aaguid=aaguid, aaguid=aaguid,
transports=transports transports=transports,
) )
passkey.create() passkey.create()
logger.info(f"用户 {current_user.name} 成功注册PassKey: {passkey_req.name}") logger.info(f"用户 {current_user.name} 成功注册PassKey: {passkey_req.name}")
return schemas.Response( return schemas.Response(success=True, message="通行密钥注册成功")
success=True,
message="通行密钥注册成功"
)
except Exception as e: except Exception as e:
logger.error(f"注册PassKey失败: {e}") logger.error(f"注册PassKey失败: {e}")
return schemas.Response( return schemas.Response(success=False, message=f"注册失败: {str(e)}")
success=False,
message=f"注册失败: {str(e)}"
@router.post(
"/passkey/authenticate/start",
summary="开始 PassKey 认证",
response_model=schemas.Response,
) )
@router.post("/passkey/authenticate/start", summary="开始 PassKey 认证", response_model=schemas.Response)
def passkey_authenticate_start( def passkey_authenticate_start(
passkey_req: PassKeyAuthenticationStart = Body(...) passkey_req: PassKeyAuthenticationStart = Body(...),
) -> Any: ) -> Any:
"""开始 PassKey 认证 - 生成认证选项""" """开始 PassKey 认证 - 生成认证选项"""
try: try:
@@ -296,13 +329,12 @@ def passkey_authenticate_start(
# 如果指定了用户名,只允许该用户的PassKey # 如果指定了用户名,只允许该用户的PassKey
if passkey_req.username: if passkey_req.username:
user = User.get_by_name(db=None, name=passkey_req.username) user = User.get_by_name(db=None, name=passkey_req.username)
existing_passkeys = PassKey.get_by_user_id(db=None, user_id=user.id) if user else None existing_passkeys = (
PassKey.get_by_user_id(db=None, user_id=user.id) if user else None
)
if not user or not existing_passkeys: if not user or not existing_passkeys:
return schemas.Response( return schemas.Response(success=False, message="认证失败")
success=False,
message="认证失败"
)
existing_credentials = _build_credential_list(existing_passkeys) existing_credentials = _build_credential_list(existing_passkeys)
@@ -312,29 +344,26 @@ def passkey_authenticate_start(
) )
return schemas.Response( return schemas.Response(
success=True, success=True, data={"options": options_json, "challenge": challenge}
data={
'options': options_json,
'challenge': challenge
}
) )
except Exception as e: except Exception as e:
logger.error(f"生成PassKey认证选项失败: {e}") logger.error(f"生成PassKey认证选项失败: {e}")
return schemas.Response( return schemas.Response(success=False, message="认证失败")
success=False,
message="认证失败"
@router.post(
"/passkey/authenticate/finish",
summary="完成 PassKey 认证",
response_model=schemas.Token,
) )
def passkey_authenticate_finish(passkey_req: PassKeyAuthenticationFinish) -> Any:
@router.post("/passkey/authenticate/finish", summary="完成 PassKey 认证", response_model=schemas.Token)
def passkey_authenticate_finish(
passkey_req: PassKeyAuthenticationFinish
) -> Any:
"""完成 PassKey 认证 - 验证凭证并返回 token""" """完成 PassKey 认证 - 验证凭证并返回 token"""
try: try:
# 提取并标准化凭证ID # 提取并标准化凭证ID
try: try:
credential_id = _extract_and_standardize_credential_id(passkey_req.credential) credential_id = _extract_and_standardize_credential_id(
passkey_req.credential
)
except ValueError as e: except ValueError as e:
logger.warning(f"PassKey认证失败,提供的凭证无效: {e}") logger.warning(f"PassKey认证失败,提供的凭证无效: {e}")
raise HTTPException(status_code=401, detail="认证失败") raise HTTPException(status_code=401, detail="认证失败")
@@ -349,7 +378,7 @@ def passkey_authenticate_finish(
success, _ = _verify_passkey_and_update( success, _ = _verify_passkey_and_update(
credential=passkey_req.credential, credential=passkey_req.credential,
challenge=passkey_req.challenge, challenge=passkey_req.challenge,
passkey=passkey passkey=passkey,
) )
if not success: if not success:
@@ -359,7 +388,10 @@ def passkey_authenticate_finish(
# 生成token # 生成token
level = SitesHelper().auth_level level = SitesHelper().auth_level
show_wizard = not SystemConfigOper().get(SystemConfigKey.SetupWizardState) and not settings.ADVANCED_MODE show_wizard = (
not SystemConfigOper().get(SystemConfigKey.SetupWizardState)
and not settings.ADVANCED_MODE
)
return schemas.Token( return schemas.Token(
access_token=security.create_access_token( access_token=security.create_access_token(
@@ -367,7 +399,7 @@ def passkey_authenticate_finish(
username=user.name, username=user.name,
super_user=user.is_superuser, super_user=user.is_superuser,
expires_delta=timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES), expires_delta=timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES),
level=level level=level,
), ),
token_type="bearer", token_type="bearer",
super_user=user.is_superuser, super_user=user.is_superuser,
@@ -376,7 +408,7 @@ def passkey_authenticate_finish(
avatar=user.avatar, avatar=user.avatar,
level=level, level=level,
permissions=user.permissions or {}, permissions=user.permissions or {},
wizard=show_wizard wizard=show_wizard,
) )
except HTTPException: except HTTPException:
raise raise
@@ -385,80 +417,83 @@ def passkey_authenticate_finish(
raise HTTPException(status_code=401, detail="认证失败") raise HTTPException(status_code=401, detail="认证失败")
@router.get("/passkey/list", summary="获取当前用户的 PassKey 列表", response_model=schemas.Response) @router.get(
"/passkey/list",
summary="获取当前用户的 PassKey 列表",
response_model=schemas.Response,
)
def passkey_list( def passkey_list(
current_user: Annotated[User, Depends(get_current_active_user)] current_user: Annotated[User, Depends(get_current_active_user)],
) -> Any: ) -> Any:
"""获取当前用户的所有 PassKey""" """获取当前用户的所有 PassKey"""
try: try:
passkeys = PassKey.get_by_user_id(db=None, user_id=current_user.id) passkeys = PassKey.get_by_user_id(db=None, user_id=current_user.id)
key_list = [ key_list = (
[
{ {
'id': pk.id, "id": pk.id,
'name': pk.name, "name": pk.name,
'created_at': pk.created_at.isoformat() if pk.created_at else None, "created_at": pk.created_at.isoformat() if pk.created_at else None,
'last_used_at': pk.last_used_at.isoformat() if pk.last_used_at else None, "last_used_at": pk.last_used_at.isoformat()
'aaguid': pk.aaguid, if pk.last_used_at
'transports': pk.transports else None,
"aaguid": pk.aaguid,
"transports": pk.transports,
} }
for pk in passkeys for pk in passkeys
] if passkeys else [] ]
if passkeys
return schemas.Response( else []
success=True,
data=key_list
) )
return schemas.Response(success=True, data=key_list)
except Exception as e: except Exception as e:
logger.error(f"获取PassKey列表失败: {e}") logger.error(f"获取PassKey列表失败: {e}")
return schemas.Response( return schemas.Response(success=False, message=f"获取列表失败: {str(e)}")
success=False,
message=f"获取列表失败: {str(e)}"
)
@router.post("/passkey/delete", summary="删除 PassKey", response_model=schemas.Response) @router.post("/passkey/delete", summary="删除 PassKey", response_model=schemas.Response)
async def passkey_delete( async def passkey_delete(
data: PassKeyDeleteRequest, data: PassKeyDeleteRequest,
current_user: User = Depends(get_current_active_user_async) current_user: User = Depends(get_current_active_user_async),
) -> Any: ) -> Any:
"""删除指定的 PassKey""" """删除指定的 PassKey"""
try: try:
# 验证密码 # 验证密码
if not security.verify_password(data.password, str(current_user.hashed_password)): if not security.verify_password(
data.password, str(current_user.hashed_password)
):
return schemas.Response(success=False, message="密码错误") return schemas.Response(success=False, message="密码错误")
success = PassKey.delete_by_id(db=None, passkey_id=data.passkey_id, user_id=current_user.id) success = PassKey.delete_by_id(
db=None, passkey_id=data.passkey_id, user_id=current_user.id
)
if success: if success:
logger.info(f"用户 {current_user.name} 删除了PassKey: {data.passkey_id}") logger.info(f"用户 {current_user.name} 删除了PassKey: {data.passkey_id}")
return schemas.Response( return schemas.Response(success=True, message="通行密钥已删除")
success=True,
message="通行密钥已删除"
)
else: else:
return schemas.Response( return schemas.Response(success=False, message="通行密钥不存在或无权删除")
success=False,
message="通行密钥不存在或无权删除"
)
except Exception as e: except Exception as e:
logger.error(f"删除PassKey失败: {e}") logger.error(f"删除PassKey失败: {e}")
return schemas.Response( return schemas.Response(success=False, message=f"删除失败: {str(e)}")
success=False,
message=f"删除失败: {str(e)}"
@router.post(
"/passkey/verify", summary="PassKey 二次验证", response_model=schemas.Response
) )
@router.post("/passkey/verify", summary="PassKey 二次验证", response_model=schemas.Response)
def passkey_verify_mfa( def passkey_verify_mfa(
passkey_req: PassKeyAuthenticationFinish, passkey_req: PassKeyAuthenticationFinish,
current_user: Annotated[User, Depends(get_current_active_user)] current_user: Annotated[User, Depends(get_current_active_user)],
) -> Any: ) -> Any:
"""使用 PassKey 进行二次验证(MFA""" """使用 PassKey 进行二次验证(MFA"""
try: try:
# 提取并标准化凭证ID # 提取并标准化凭证ID
try: try:
credential_id = _extract_and_standardize_credential_id(passkey_req.credential) credential_id = _extract_and_standardize_credential_id(
passkey_req.credential
)
except ValueError as e: except ValueError as e:
logger.warning(f"PassKey二次验证失败,提供的凭证无效: {e}") logger.warning(f"PassKey二次验证失败,提供的凭证无效: {e}")
return schemas.Response(success=False, message="验证失败") return schemas.Response(success=False, message="验证失败")
@@ -467,32 +502,22 @@ def passkey_verify_mfa(
passkey = PassKey.get_by_credential_id(db=None, credential_id=credential_id) passkey = PassKey.get_by_credential_id(db=None, credential_id=credential_id)
if not passkey or passkey.user_id != current_user.id: if not passkey or passkey.user_id != current_user.id:
return schemas.Response( return schemas.Response(
success=False, success=False, message="通行密钥不存在或不属于当前用户"
message="通行密钥不存在或不属于当前用户"
) )
# 验证认证响应并更新 # 验证认证响应并更新
success, _ = _verify_passkey_and_update( success, _ = _verify_passkey_and_update(
credential=passkey_req.credential, credential=passkey_req.credential,
challenge=passkey_req.challenge, challenge=passkey_req.challenge,
passkey=passkey passkey=passkey,
) )
if not success: if not success:
return schemas.Response( return schemas.Response(success=False, message="通行密钥验证失败")
success=False,
message="通行密钥验证失败"
)
logger.info(f"用户 {current_user.name} 通过PassKey二次验证成功") logger.info(f"用户 {current_user.name} 通过PassKey二次验证成功")
return schemas.Response( return schemas.Response(success=True, message="二次验证成功")
success=True,
message="二次验证成功"
)
except Exception as e: except Exception as e:
logger.error(f"PassKey二次验证失败: {e}") logger.error(f"PassKey二次验证失败: {e}")
return schemas.Response( return schemas.Response(success=False, message="验证失败")
success=False,
message="验证失败"
)
+28 -8
View File
@@ -251,9 +251,15 @@ def _check_auth(
return None return None
@router.get("/models", summary="OpenAI compatible models", response_model=schemas.OpenAIModelListResponse) @router.get(
"/models",
summary="OpenAI compatible models",
response_model=schemas.OpenAIModelListResponse,
)
async def list_models( async def list_models(
credentials: Optional[HTTPAuthorizationCredentials] = Security(openai_bearer_scheme), credentials: Optional[HTTPAuthorizationCredentials] = Security(
openai_bearer_scheme
),
): ):
auth_error = _check_auth(credentials) auth_error = _check_auth(credentials)
if auth_error: if auth_error:
@@ -272,7 +278,9 @@ async def list_models(
async def chat_completions( async def chat_completions(
payload: schemas.OpenAIChatCompletionsRequest, payload: schemas.OpenAIChatCompletionsRequest,
request: Request, request: Request,
credentials: Optional[HTTPAuthorizationCredentials] = Security(openai_bearer_scheme), credentials: Optional[HTTPAuthorizationCredentials] = Security(
openai_bearer_scheme
),
): ):
auth_error = _check_auth(credentials) auth_error = _check_auth(credentials)
if auth_error: if auth_error:
@@ -304,7 +312,9 @@ async def chat_completions(
) )
try: try:
prompt, images = build_prompt(payload.messages, use_server_session=use_server_session) prompt, images = build_prompt(
payload.messages, use_server_session=use_server_session
)
except ValueError as exc: except ValueError as exc:
return _error_response(str(exc), 400, code="invalid_messages") return _error_response(str(exc), 400, code="invalid_messages")
@@ -353,10 +363,16 @@ async def chat_completions(
return JSONResponse(content=build_completion_payload(content, MODEL_ID)) return JSONResponse(content=build_completion_payload(content, MODEL_ID))
@router.post("/responses", summary="OpenAI compatible responses", response_model=schemas.OpenAIResponsesResponse) @router.post(
"/responses",
summary="OpenAI compatible responses",
response_model=schemas.OpenAIResponsesResponse,
)
async def responses( async def responses(
payload: schemas.OpenAIResponsesRequest, payload: schemas.OpenAIResponsesRequest,
credentials: Optional[HTTPAuthorizationCredentials] = Security(openai_bearer_scheme), credentials: Optional[HTTPAuthorizationCredentials] = Security(
openai_bearer_scheme
),
): ):
auth_error = _check_auth(credentials) auth_error = _check_auth(credentials)
if auth_error: if auth_error:
@@ -377,7 +393,9 @@ async def responses(
code="unsupported_stream", code="unsupported_stream",
) )
normalized_messages = build_responses_input(payload.input, instructions=payload.instructions) normalized_messages = build_responses_input(
payload.input, instructions=payload.instructions
)
if not normalized_messages: if not normalized_messages:
return _error_response( return _error_response(
"`input` must include at least one usable message.", "`input` must include at least one usable message.",
@@ -386,7 +404,9 @@ async def responses(
) )
try: try:
prompt, images = build_prompt(normalized_messages, use_server_session=bool(payload.user)) prompt, images = build_prompt(
normalized_messages, use_server_session=bool(payload.user)
)
except ValueError as exc: except ValueError as exc:
return _error_response(str(exc), 400, code="invalid_input") return _error_response(str(exc), 400, code="invalid_input")
+178 -85
View File
@@ -16,7 +16,10 @@ from app.core.plugin import PluginManager
from app.core.security import verify_apikey, verify_token from app.core.security import verify_apikey, verify_token
from app.db.models import User from app.db.models import User
from app.db.systemconfig_oper import SystemConfigOper from app.db.systemconfig_oper import SystemConfigOper
from app.db.user_oper import get_current_active_superuser, get_current_active_superuser_async from app.db.user_oper import (
get_current_active_superuser,
get_current_active_superuser_async,
)
from app.factory import app from app.factory import app
from app.helper.plugin import PluginHelper from app.helper.plugin import PluginHelper
from app.log import logger from app.log import logger
@@ -76,7 +79,10 @@ def _update_plugin_api_routes(plugin_id: Optional[str], action: str):
auth_mode = api.pop("auth", "apikey") auth_mode = api.pop("auth", "apikey")
dependencies = api.setdefault("dependencies", []) dependencies = api.setdefault("dependencies", [])
if not allow_anonymous: if not allow_anonymous:
if auth_mode == "bear" and Depends(verify_token) not in dependencies: if (
auth_mode == "bear"
and Depends(verify_token) not in dependencies
):
dependencies.append(Depends(verify_token)) dependencies.append(Depends(verify_token))
elif Depends(verify_apikey) not in dependencies: elif Depends(verify_apikey) not in dependencies:
dependencies.append(Depends(verify_apikey)) dependencies.append(Depends(verify_apikey))
@@ -140,8 +146,11 @@ def register_plugin(plugin_id: str):
@router.get("/", summary="所有插件", response_model=List[schemas.Plugin]) @router.get("/", summary="所有插件", response_model=List[schemas.Plugin])
async def all_plugins(_: User = Depends(get_current_active_superuser_async), async def all_plugins(
state: Optional[str] = "all", force: bool = False) -> List[schemas.Plugin]: _: User = Depends(get_current_active_superuser_async),
state: Optional[str] = "all",
force: bool = False,
) -> List[schemas.Plugin]:
""" """
查询所有插件清单,包括本地插件和在线插件,插件状态:installed, market, all 查询所有插件清单,包括本地插件和在线插件,插件状态:installed, market, all
""" """
@@ -159,8 +168,11 @@ async def all_plugins(_: User = Depends(get_current_active_superuser_async),
local_repo_plugins = plugin_manager.get_local_repo_plugins() local_repo_plugins = plugin_manager.get_local_repo_plugins()
# 在线插件 # 在线插件
online_plugins = await plugin_manager.async_get_online_plugins(force) online_plugins = await plugin_manager.async_get_online_plugins(force)
candidate_plugins = plugin_manager.process_plugins_list(online_plugins + local_repo_plugins, []) \ candidate_plugins = (
if online_plugins or local_repo_plugins else [] plugin_manager.process_plugins_list(online_plugins + local_repo_plugins, [])
if online_plugins or local_repo_plugins
else []
)
if not candidate_plugins: if not candidate_plugins:
# 没有获取在线插件 # 没有获取在线插件
if state == "market": if state == "market":
@@ -208,8 +220,12 @@ async def statistic(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
return await PluginHelper().async_get_statistic() return await PluginHelper().async_get_statistic()
@router.get("/reload/{plugin_id}", summary="重新加载插件", response_model=schemas.Response) @router.get(
def reload_plugin(plugin_id: str, _: User = Depends(get_current_active_superuser)) -> Any: "/reload/{plugin_id}", summary="重新加载插件", response_model=schemas.Response
)
def reload_plugin(
plugin_id: str, _: User = Depends(get_current_active_superuser)
) -> Any:
""" """
重新加载插件 重新加载插件
""" """
@@ -221,10 +237,12 @@ def reload_plugin(plugin_id: str, _: User = Depends(get_current_active_superuser
@router.get("/install/{plugin_id}", summary="安装插件", response_model=schemas.Response) @router.get("/install/{plugin_id}", summary="安装插件", response_model=schemas.Response)
async def install(plugin_id: str, async def install(
plugin_id: str,
repo_url: Optional[str] = "", repo_url: Optional[str] = "",
force: Optional[bool] = False, force: Optional[bool] = False,
_: User = Depends(get_current_active_superuser_async)) -> Any: _: User = Depends(get_current_active_superuser_async),
) -> Any:
""" """
安装插件 安装插件
""" """
@@ -238,21 +256,23 @@ async def install(plugin_id: str,
# 插件不存在或需要强制安装,下载安装并注册插件 # 插件不存在或需要强制安装,下载安装并注册插件
if repo_url: if repo_url:
state, msg = await plugin_helper.async_install( state, msg = await plugin_helper.async_install(
pid=plugin_id, pid=plugin_id, repo_url=repo_url, force_install=force
repo_url=repo_url,
force_install=force
) )
# 安装失败则直接响应 # 安装失败则直接响应
if not state: if not state:
return schemas.Response(success=False, message=msg) return schemas.Response(success=False, message=msg)
else: else:
# repo_url 为空时,也直接响应 # repo_url 为空时,也直接响应
return schemas.Response(success=False, message="没有传入仓库地址,无法正确安装插件,请检查配置") return schemas.Response(
success=False, message="没有传入仓库地址,无法正确安装插件,请检查配置"
)
# 安装插件 # 安装插件
if plugin_id not in install_plugins: if plugin_id not in install_plugins:
install_plugins.append(plugin_id) install_plugins.append(plugin_id)
# 保存设置 # 保存设置
await SystemConfigOper().async_set(SystemConfigKey.UserInstalledPlugins, install_plugins) await SystemConfigOper().async_set(
SystemConfigKey.UserInstalledPlugins, install_plugins
)
# 重新加载插件 # 重新加载插件
await run_in_threadpool(reload_plugin, plugin_id) await run_in_threadpool(reload_plugin, plugin_id)
return schemas.Response(success=True) return schemas.Response(success=True)
@@ -268,7 +288,11 @@ async def remotes(token: str) -> Any:
return PluginManager().get_plugin_remotes() return PluginManager().get_plugin_remotes()
@router.get("/sidebar_nav", summary="获取插件侧栏导航项", response_model=List[schemas.PluginSidebarNavItem]) @router.get(
"/sidebar_nav",
summary="获取插件侧栏导航项",
response_model=List[schemas.PluginSidebarNavItem],
)
def plugin_sidebar_nav(_: schemas.TokenPayload = Depends(verify_token)) -> Any: def plugin_sidebar_nav(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
""" """
聚合已启用 Vue 插件声明的侧栏入口(get_sidebar_nav),供前端主界面侧栏展示。 聚合已启用 Vue 插件声明的侧栏入口(get_sidebar_nav),供前端主界面侧栏展示。
@@ -277,15 +301,19 @@ def plugin_sidebar_nav(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
@router.get("/form/{plugin_id}", summary="获取插件表单页面") @router.get("/form/{plugin_id}", summary="获取插件表单页面")
def plugin_form(plugin_id: str, def plugin_form(
_: User = Depends(get_current_active_superuser)) -> dict: plugin_id: str, _: User = Depends(get_current_active_superuser)
) -> dict:
""" """
根据插件ID获取插件配置表单或Vue组件URL 根据插件ID获取插件配置表单或Vue组件URL
""" """
plugin_manager = PluginManager() plugin_manager = PluginManager()
plugin_instance = plugin_manager.running_plugins.get(plugin_id) plugin_instance = plugin_manager.running_plugins.get(plugin_id)
if not plugin_instance: if not plugin_instance:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=f"插件 {plugin_id} 不存在或未加载") raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"插件 {plugin_id} 不存在或未加载",
)
# 渲染模式 # 渲染模式
render_mode, _ = plugin_instance.get_render_mode() render_mode, _ = plugin_instance.get_render_mode()
@@ -294,7 +322,7 @@ def plugin_form(plugin_id: str,
return { return {
"render_mode": render_mode, "render_mode": render_mode,
"conf": conf, "conf": conf,
"model": plugin_manager.get_plugin_config(plugin_id) or model "model": plugin_manager.get_plugin_config(plugin_id) or model,
} }
except Exception as e: except Exception as e:
logger.error(f"插件 {plugin_id} 调用方法 get_form 出错: {str(e)}") logger.error(f"插件 {plugin_id} 调用方法 get_form 出错: {str(e)}")
@@ -302,29 +330,33 @@ def plugin_form(plugin_id: str,
@router.get("/page/{plugin_id}", summary="获取插件数据页面") @router.get("/page/{plugin_id}", summary="获取插件数据页面")
def plugin_page(plugin_id: str, _: User = Depends(get_current_active_superuser)) -> dict: def plugin_page(
plugin_id: str, _: User = Depends(get_current_active_superuser)
) -> dict:
""" """
根据插件ID获取插件数据页面 根据插件ID获取插件数据页面
""" """
plugin_instance = PluginManager().running_plugins.get(plugin_id) plugin_instance = PluginManager().running_plugins.get(plugin_id)
if not plugin_instance: if not plugin_instance:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=f"插件 {plugin_id} 不存在或未加载") raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"插件 {plugin_id} 不存在或未加载",
)
# 渲染模式 # 渲染模式
render_mode, _ = plugin_instance.get_render_mode() render_mode, _ = plugin_instance.get_render_mode()
try: try:
page = plugin_instance.get_page() page = plugin_instance.get_page()
return { return {"render_mode": render_mode, "page": page or []}
"render_mode": render_mode,
"page": page or []
}
except Exception as e: except Exception as e:
logger.error(f"插件 {plugin_id} 调用方法 get_page 出错: {str(e)}") logger.error(f"插件 {plugin_id} 调用方法 get_page 出错: {str(e)}")
return {} return {}
@router.get("/dashboard/meta", summary="获取所有插件仪表板元信息") @router.get("/dashboard/meta", summary="获取所有插件仪表板元信息")
def plugin_dashboard_meta(_: schemas.TokenPayload = Depends(verify_token)) -> List[dict]: def plugin_dashboard_meta(
_: schemas.TokenPayload = Depends(verify_token),
) -> List[dict]:
""" """
获取所有插件仪表板元信息 获取所有插件仪表板元信息
""" """
@@ -332,8 +364,12 @@ def plugin_dashboard_meta(_: schemas.TokenPayload = Depends(verify_token)) -> Li
@router.get("/dashboard/{plugin_id}/{key}", summary="获取插件仪表板配置") @router.get("/dashboard/{plugin_id}/{key}", summary="获取插件仪表板配置")
def plugin_dashboard_by_key(plugin_id: str, key: str, user_agent: Annotated[str | None, Header()] = None, def plugin_dashboard_by_key(
_: schemas.TokenPayload = Depends(verify_token)) -> Optional[schemas.PluginDashboard]: plugin_id: str,
key: str,
user_agent: Annotated[str | None, Header()] = None,
_: schemas.TokenPayload = Depends(verify_token),
) -> Optional[schemas.PluginDashboard]:
""" """
根据插件ID获取插件仪表板 根据插件ID获取插件仪表板
""" """
@@ -341,17 +377,23 @@ def plugin_dashboard_by_key(plugin_id: str, key: str, user_agent: Annotated[str
@router.get("/dashboard/{plugin_id}", summary="获取插件仪表板配置") @router.get("/dashboard/{plugin_id}", summary="获取插件仪表板配置")
def plugin_dashboard(plugin_id: str, user_agent: Annotated[str | None, Header()] = None, def plugin_dashboard(
_: schemas.TokenPayload = Depends(verify_token)) -> schemas.PluginDashboard: plugin_id: str,
user_agent: Annotated[str | None, Header()] = None,
_: schemas.TokenPayload = Depends(verify_token),
) -> schemas.PluginDashboard:
""" """
根据插件ID获取插件仪表板 根据插件ID获取插件仪表板
""" """
return plugin_dashboard_by_key(plugin_id, "", user_agent) return plugin_dashboard_by_key(plugin_id, "", user_agent)
@router.get("/reset/{plugin_id}", summary="重置插件配置及数据", response_model=schemas.Response) @router.get(
def reset_plugin(plugin_id: str, "/reset/{plugin_id}", summary="重置插件配置及数据", response_model=schemas.Response
_: User = Depends(get_current_active_superuser)) -> Any: )
def reset_plugin(
plugin_id: str, _: User = Depends(get_current_active_superuser)
) -> Any:
""" """
根据插件ID重置插件配置及数据 根据插件ID重置插件配置及数据
""" """
@@ -372,42 +414,54 @@ async def plugin_static_file(plugin_id: str, filepath: str):
""" """
# 基础安全检查 # 基础安全检查
if ".." in filepath or ".." in plugin_id: if ".." in filepath or ".." in plugin_id:
logger.warning(f"Static File API: Path traversal attempt detected: {plugin_id}/{filepath}") logger.warning(
f"Static File API: Path traversal attempt detected: {plugin_id}/{filepath}"
)
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Forbidden") raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Forbidden")
plugin_base_dir = AsyncPath(settings.ROOT_PATH) / "app" / "plugins" / plugin_id.lower() plugin_base_dir = (
plugin_file_path = plugin_base_dir / filepath.lstrip('/') AsyncPath(settings.ROOT_PATH) / "app" / "plugins" / plugin_id.lower()
)
plugin_file_path = plugin_base_dir / filepath.lstrip("/")
try: try:
resolved_base = await plugin_base_dir.resolve() resolved_base = await plugin_base_dir.resolve()
resolved_file = await plugin_file_path.resolve() resolved_file = await plugin_file_path.resolve()
except Exception: except Exception:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid path") raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid path"
)
if not resolved_file.is_relative_to(resolved_base): if not resolved_file.is_relative_to(resolved_base):
logger.warning(f"Static File API: Path traversal attempt detected: {plugin_id}/{filepath}") logger.warning(
f"Static File API: Path traversal attempt detected: {plugin_id}/{filepath}"
)
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Forbidden") raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Forbidden")
if not await plugin_file_path.exists(): if not await plugin_file_path.exists():
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=f"{plugin_file_path} 不存在") raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND, detail=f"{plugin_file_path} 不存在"
)
if not await plugin_file_path.is_file(): if not await plugin_file_path.is_file():
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=f"{plugin_file_path} 不是文件") raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN, detail=f"{plugin_file_path} 不是文件"
)
# 判断 MIME 类型 # 判断 MIME 类型
response_type, _ = mimetypes.guess_type(str(plugin_file_path)) response_type, _ = mimetypes.guess_type(str(plugin_file_path))
suffix = plugin_file_path.suffix.lower() suffix = plugin_file_path.suffix.lower()
# 强制修正 .mjs 和 .js 的 MIME 类型 # 强制修正 .mjs 和 .js 的 MIME 类型
if suffix in ['.js', '.mjs']: if suffix in [".js", ".mjs"]:
response_type = 'application/javascript' response_type = "application/javascript"
elif suffix == '.css' and not response_type: # 如果 guess_type 没猜对 css,也修正 elif suffix == ".css" and not response_type: # 如果 guess_type 没猜对 css,也修正
response_type = 'text/css' response_type = "text/css"
elif not response_type: # 对于其他猜不出的类型 elif not response_type: # 对于其他猜不出的类型
response_type = 'application/octet-stream' response_type = "application/octet-stream"
try: try:
# 异步生成器函数,用于流式读取文件 # 异步生成器函数,用于流式读取文件
async def file_generator(): async def file_generator():
async with aiofiles.open(plugin_file_path, mode='rb') as file: async with aiofiles.open(plugin_file_path, mode="rb") as file:
# 8KB 块大小 # 8KB 块大小
while chunk := await file.read(8192): while chunk := await file.read(8192):
yield chunk yield chunk
@@ -415,15 +469,22 @@ async def plugin_static_file(plugin_id: str, filepath: str):
return StreamingResponse( return StreamingResponse(
file_generator(), file_generator(),
media_type=response_type, media_type=response_type,
headers={"Content-Disposition": f"inline; filename={plugin_file_path.name}"} headers={
"Content-Disposition": f"inline; filename={plugin_file_path.name}"
},
) )
except Exception as e: except Exception as e:
logger.error(f"Error creating/sending StreamingResponse for {plugin_file_path}: {e}", exc_info=True) logger.error(
f"Error creating/sending StreamingResponse for {plugin_file_path}: {e}",
exc_info=True,
)
raise HTTPException(status_code=500, detail="Internal Server Error") raise HTTPException(status_code=500, detail="Internal Server Error")
@router.get("/folders", summary="获取插件文件夹配置", response_model=dict) @router.get("/folders", summary="获取插件文件夹配置", response_model=dict)
async def get_plugin_folders(_: User = Depends(get_current_active_superuser_async)) -> dict: async def get_plugin_folders(
_: User = Depends(get_current_active_superuser_async),
) -> dict:
""" """
获取插件文件夹分组配置 获取插件文件夹分组配置
""" """
@@ -436,7 +497,9 @@ async def get_plugin_folders(_: User = Depends(get_current_active_superuser_asyn
@router.post("/folders", summary="保存插件文件夹配置", response_model=schemas.Response) @router.post("/folders", summary="保存插件文件夹配置", response_model=schemas.Response)
async def save_plugin_folders(folders: dict, _: User = Depends(get_current_active_superuser_async)) -> Any: async def save_plugin_folders(
folders: dict, _: User = Depends(get_current_active_superuser_async)
) -> Any:
""" """
保存插件文件夹分组配置 保存插件文件夹分组配置
""" """
@@ -448,9 +511,12 @@ async def save_plugin_folders(folders: dict, _: User = Depends(get_current_activ
return schemas.Response(success=False, message=str(e)) return schemas.Response(success=False, message=str(e))
@router.post("/folders/{folder_name}", summary="创建插件文件夹", response_model=schemas.Response) @router.post(
async def create_plugin_folder(folder_name: str, "/folders/{folder_name}", summary="创建插件文件夹", response_model=schemas.Response
_: User = Depends(get_current_active_superuser_async)) -> Any: )
async def create_plugin_folder(
folder_name: str, _: User = Depends(get_current_active_superuser_async)
) -> Any:
""" """
创建新的插件文件夹 创建新的插件文件夹
""" """
@@ -458,14 +524,19 @@ async def create_plugin_folder(folder_name: str,
if folder_name not in folders: if folder_name not in folders:
folders[folder_name] = [] folders[folder_name] = []
SystemConfigOper().set(SystemConfigKey.PluginFolders, folders) SystemConfigOper().set(SystemConfigKey.PluginFolders, folders)
return schemas.Response(success=True, message=f"文件夹 '{folder_name}' 创建成功") return schemas.Response(
success=True, message=f"文件夹 '{folder_name}' 创建成功"
)
else: else:
return schemas.Response(success=False, message=f"文件夹 '{folder_name}' 已存在") return schemas.Response(success=False, message=f"文件夹 '{folder_name}' 已存在")
@router.delete("/folders/{folder_name}", summary="删除插件文件夹", response_model=schemas.Response) @router.delete(
async def delete_plugin_folder(folder_name: str, "/folders/{folder_name}", summary="删除插件文件夹", response_model=schemas.Response
_: User = Depends(get_current_active_superuser_async)) -> Any: )
async def delete_plugin_folder(
folder_name: str, _: User = Depends(get_current_active_superuser_async)
) -> Any:
""" """
删除插件文件夹 删除插件文件夹
""" """
@@ -473,27 +544,40 @@ async def delete_plugin_folder(folder_name: str,
if folder_name in folders: if folder_name in folders:
del folders[folder_name] del folders[folder_name]
await SystemConfigOper().async_set(SystemConfigKey.PluginFolders, folders) await SystemConfigOper().async_set(SystemConfigKey.PluginFolders, folders)
return schemas.Response(success=True, message=f"文件夹 '{folder_name}' 删除成功") return schemas.Response(
success=True, message=f"文件夹 '{folder_name}' 删除成功"
)
else: else:
return schemas.Response(success=False, message=f"文件夹 '{folder_name}' 不存在") return schemas.Response(success=False, message=f"文件夹 '{folder_name}' 不存在")
@router.put("/folders/{folder_name}/plugins", summary="更新文件夹中的插件", response_model=schemas.Response) @router.put(
async def update_folder_plugins(folder_name: str, plugin_ids: List[str], "/folders/{folder_name}/plugins",
_: User = Depends(get_current_active_superuser_async)) -> Any: summary="更新文件夹中的插件",
response_model=schemas.Response,
)
async def update_folder_plugins(
folder_name: str,
plugin_ids: List[str],
_: User = Depends(get_current_active_superuser_async),
) -> Any:
""" """
更新指定文件夹中的插件列表 更新指定文件夹中的插件列表
""" """
folders = SystemConfigOper().get(SystemConfigKey.PluginFolders) or {} folders = SystemConfigOper().get(SystemConfigKey.PluginFolders) or {}
folders[folder_name] = plugin_ids folders[folder_name] = plugin_ids
await SystemConfigOper().async_set(SystemConfigKey.PluginFolders, folders) await SystemConfigOper().async_set(SystemConfigKey.PluginFolders, folders)
return schemas.Response(success=True, message=f"文件夹 '{folder_name}' 中的插件已更新") return schemas.Response(
success=True, message=f"文件夹 '{folder_name}' 中的插件已更新"
)
@router.post("/clone/{plugin_id}", summary="创建插件分身", response_model=schemas.Response) @router.post(
def clone_plugin(plugin_id: str, "/clone/{plugin_id}", summary="创建插件分身", response_model=schemas.Response
clone_data: dict, )
_: User = Depends(get_current_active_superuser)) -> Any: def clone_plugin(
plugin_id: str, clone_data: dict, _: User = Depends(get_current_active_superuser)
) -> Any:
""" """
创建插件分身 创建插件分身
""" """
@@ -504,7 +588,7 @@ def clone_plugin(plugin_id: str,
name=clone_data.get("name", ""), name=clone_data.get("name", ""),
description=clone_data.get("description", ""), description=clone_data.get("description", ""),
version=clone_data.get("version", ""), version=clone_data.get("version", ""),
icon=clone_data.get("icon", "") icon=clone_data.get("icon", ""),
) )
if success: if success:
@@ -521,8 +605,9 @@ def clone_plugin(plugin_id: str,
@router.get("/{plugin_id}", summary="获取插件配置") @router.get("/{plugin_id}", summary="获取插件配置")
async def plugin_config(plugin_id: str, async def plugin_config(
_: User = Depends(get_current_active_superuser_async)) -> dict: plugin_id: str, _: User = Depends(get_current_active_superuser_async)
) -> dict:
""" """
根据插件ID获取插件配置信息 根据插件ID获取插件配置信息
""" """
@@ -530,8 +615,9 @@ async def plugin_config(plugin_id: str,
@router.put("/{plugin_id}", summary="更新插件配置", response_model=schemas.Response) @router.put("/{plugin_id}", summary="更新插件配置", response_model=schemas.Response)
def set_plugin_config(plugin_id: str, conf: dict, def set_plugin_config(
_: User = Depends(get_current_active_superuser)) -> Any: plugin_id: str, conf: dict, _: User = Depends(get_current_active_superuser)
) -> Any:
""" """
更新插件配置 更新插件配置
""" """
@@ -546,8 +632,9 @@ def set_plugin_config(plugin_id: str, conf: dict,
@router.delete("/{plugin_id}", summary="卸载插件", response_model=schemas.Response) @router.delete("/{plugin_id}", summary="卸载插件", response_model=schemas.Response)
def uninstall_plugin(plugin_id: str, def uninstall_plugin(
_: User = Depends(get_current_active_superuser)) -> Any: plugin_id: str, _: User = Depends(get_current_active_superuser)
) -> Any:
""" """
卸载插件 卸载插件
""" """
@@ -599,9 +686,9 @@ def _add_clone_to_plugin_folder(original_plugin_id: str, clone_plugin_id: str):
# 查找原插件所在的文件夹 # 查找原插件所在的文件夹
target_folder = None target_folder = None
for folder_name, folder_data in folders.items(): for folder_name, folder_data in folders.items():
if isinstance(folder_data, dict) and 'plugins' in folder_data: if isinstance(folder_data, dict) and "plugins" in folder_data:
# 新格式:{"plugins": [...], "order": ..., "icon": ...} # 新格式:{"plugins": [...], "order": ..., "icon": ...}
if original_plugin_id in folder_data['plugins']: if original_plugin_id in folder_data["plugins"]:
target_folder = folder_name target_folder = folder_name
break break
elif isinstance(folder_data, list): elif isinstance(folder_data, list):
@@ -613,21 +700,27 @@ def _add_clone_to_plugin_folder(original_plugin_id: str, clone_plugin_id: str):
# 如果找到了原插件所在的文件夹,则将分身插件也添加到该文件夹中 # 如果找到了原插件所在的文件夹,则将分身插件也添加到该文件夹中
if target_folder: if target_folder:
folder_data = folders[target_folder] folder_data = folders[target_folder]
if isinstance(folder_data, dict) and 'plugins' in folder_data: if isinstance(folder_data, dict) and "plugins" in folder_data:
# 新格式 # 新格式
if clone_plugin_id not in folder_data['plugins']: if clone_plugin_id not in folder_data["plugins"]:
folder_data['plugins'].append(clone_plugin_id) folder_data["plugins"].append(clone_plugin_id)
logger.info(f"已将分身插件 {clone_plugin_id} 添加到文件夹 '{target_folder}'") logger.info(
f"已将分身插件 {clone_plugin_id} 添加到文件夹 '{target_folder}'"
)
elif isinstance(folder_data, list): elif isinstance(folder_data, list):
# 旧格式 # 旧格式
if clone_plugin_id not in folder_data: if clone_plugin_id not in folder_data:
folder_data.append(clone_plugin_id) folder_data.append(clone_plugin_id)
logger.info(f"已将分身插件 {clone_plugin_id} 添加到文件夹 '{target_folder}'") logger.info(
f"已将分身插件 {clone_plugin_id} 添加到文件夹 '{target_folder}'"
)
# 保存更新后的文件夹配置 # 保存更新后的文件夹配置
config_oper.set(SystemConfigKey.PluginFolders, folders) config_oper.set(SystemConfigKey.PluginFolders, folders)
else: else:
logger.info(f"原插件 {original_plugin_id} 不在任何文件夹中,分身插件 {clone_plugin_id} 将保持独立") logger.info(
f"原插件 {original_plugin_id} 不在任何文件夹中,分身插件 {clone_plugin_id} 将保持独立"
)
except Exception as e: except Exception as e:
logger.error(f"处理插件文件夹时出错:{str(e)}") logger.error(f"处理插件文件夹时出错:{str(e)}")
@@ -649,10 +742,10 @@ def _remove_plugin_from_folders(plugin_id: str):
# 遍历所有文件夹,移除指定插件 # 遍历所有文件夹,移除指定插件
for folder_name, folder_data in folders.items(): for folder_name, folder_data in folders.items():
if isinstance(folder_data, dict) and 'plugins' in folder_data: if isinstance(folder_data, dict) and "plugins" in folder_data:
# 新格式:{"plugins": [...], "order": ..., "icon": ...} # 新格式:{"plugins": [...], "order": ..., "icon": ...}
if plugin_id in folder_data['plugins']: if plugin_id in folder_data["plugins"]:
folder_data['plugins'].remove(plugin_id) folder_data["plugins"].remove(plugin_id)
logger.info(f"已从文件夹 '{folder_name}' 中移除插件 {plugin_id}") logger.info(f"已从文件夹 '{folder_name}' 中移除插件 {plugin_id}")
modified = True modified = True
elif isinstance(folder_data, list): elif isinstance(folder_data, list):
+110 -43
View File
@@ -12,7 +12,11 @@ from app.schemas.types import ChainEventType
router = APIRouter() router = APIRouter()
@router.get("/source", summary="获取推荐数据源", response_model=List[schemas.RecommendMediaSource]) @router.get(
"/source",
summary="获取推荐数据源",
response_model=List[schemas.RecommendMediaSource],
)
def source(_: schemas.TokenPayload = Depends(verify_token)) -> Any: def source(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
""" """
获取推荐数据源 获取推荐数据源
@@ -28,104 +32,156 @@ def source(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
return [] return []
@router.get("/bangumi_calendar", summary="Bangumi每日放送", response_model=List[schemas.MediaInfo]) @router.get(
async def bangumi_calendar(page: Optional[int] = 1, "/bangumi_calendar",
summary="Bangumi每日放送",
response_model=List[schemas.MediaInfo],
)
async def bangumi_calendar(
page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
浏览Bangumi每日放送 浏览Bangumi每日放送
""" """
return await RecommendChain().async_bangumi_calendar(page=page, count=count) return await RecommendChain().async_bangumi_calendar(page=page, count=count)
@router.get("/douban_showing", summary="豆瓣正在热映", response_model=List[schemas.MediaInfo]) @router.get(
async def douban_showing(page: Optional[int] = 1, "/douban_showing", summary="豆瓣正在热映", response_model=List[schemas.MediaInfo]
)
async def douban_showing(
page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
浏览豆瓣正在热映 浏览豆瓣正在热映
""" """
return await RecommendChain().async_douban_movie_showing(page=page, count=count) return await RecommendChain().async_douban_movie_showing(page=page, count=count)
@router.get("/douban_movies", summary="豆瓣电影", response_model=List[schemas.MediaInfo]) @router.get(
async def douban_movies(sort: Optional[str] = "R", "/douban_movies", summary="豆瓣电影", response_model=List[schemas.MediaInfo]
)
async def douban_movies(
sort: Optional[str] = "R",
tags: Optional[str] = "", tags: Optional[str] = "",
page: Optional[int] = 1, page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
浏览豆瓣电影信息 浏览豆瓣电影信息
""" """
return await RecommendChain().async_douban_movies(sort=sort, tags=tags, page=page, count=count) return await RecommendChain().async_douban_movies(
sort=sort, tags=tags, page=page, count=count
)
@router.get("/douban_tvs", summary="豆瓣剧集", response_model=List[schemas.MediaInfo]) @router.get("/douban_tvs", summary="豆瓣剧集", response_model=List[schemas.MediaInfo])
async def douban_tvs(sort: Optional[str] = "R", async def douban_tvs(
sort: Optional[str] = "R",
tags: Optional[str] = "", tags: Optional[str] = "",
page: Optional[int] = 1, page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
浏览豆瓣剧集信息 浏览豆瓣剧集信息
""" """
return await RecommendChain().async_douban_tvs(sort=sort, tags=tags, page=page, count=count) return await RecommendChain().async_douban_tvs(
sort=sort, tags=tags, page=page, count=count
)
@router.get("/douban_movie_top250", summary="豆瓣电影TOP250", response_model=List[schemas.MediaInfo]) @router.get(
async def douban_movie_top250(page: Optional[int] = 1, "/douban_movie_top250",
summary="豆瓣电影TOP250",
response_model=List[schemas.MediaInfo],
)
async def douban_movie_top250(
page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
浏览豆瓣剧集信息 浏览豆瓣剧集信息
""" """
return await RecommendChain().async_douban_movie_top250(page=page, count=count) return await RecommendChain().async_douban_movie_top250(page=page, count=count)
@router.get("/douban_tv_weekly_chinese", summary="豆瓣国产剧集周榜", response_model=List[schemas.MediaInfo]) @router.get(
async def douban_tv_weekly_chinese(page: Optional[int] = 1, "/douban_tv_weekly_chinese",
summary="豆瓣国产剧集周榜",
response_model=List[schemas.MediaInfo],
)
async def douban_tv_weekly_chinese(
page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
中国每周剧集口碑榜 中国每周剧集口碑榜
""" """
return await RecommendChain().async_douban_tv_weekly_chinese(page=page, count=count) return await RecommendChain().async_douban_tv_weekly_chinese(page=page, count=count)
@router.get("/douban_tv_weekly_global", summary="豆瓣全球剧集周榜", response_model=List[schemas.MediaInfo]) @router.get(
async def douban_tv_weekly_global(page: Optional[int] = 1, "/douban_tv_weekly_global",
summary="豆瓣全球剧集周榜",
response_model=List[schemas.MediaInfo],
)
async def douban_tv_weekly_global(
page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
全球每周剧集口碑榜 全球每周剧集口碑榜
""" """
return await RecommendChain().async_douban_tv_weekly_global(page=page, count=count) return await RecommendChain().async_douban_tv_weekly_global(page=page, count=count)
@router.get("/douban_tv_animation", summary="豆瓣动画剧集", response_model=List[schemas.MediaInfo]) @router.get(
async def douban_tv_animation(page: Optional[int] = 1, "/douban_tv_animation",
summary="豆瓣动画剧集",
response_model=List[schemas.MediaInfo],
)
async def douban_tv_animation(
page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
热门动画剧集 热门动画剧集
""" """
return await RecommendChain().async_douban_tv_animation(page=page, count=count) return await RecommendChain().async_douban_tv_animation(page=page, count=count)
@router.get("/douban_movie_hot", summary="豆瓣热门电影", response_model=List[schemas.MediaInfo]) @router.get(
async def douban_movie_hot(page: Optional[int] = 1, "/douban_movie_hot", summary="豆瓣热门电影", response_model=List[schemas.MediaInfo]
)
async def douban_movie_hot(
page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
热门电影 热门电影
""" """
return await RecommendChain().async_douban_movie_hot(page=page, count=count) return await RecommendChain().async_douban_movie_hot(page=page, count=count)
@router.get("/douban_tv_hot", summary="豆瓣热门电视剧", response_model=List[schemas.MediaInfo]) @router.get(
async def douban_tv_hot(page: Optional[int] = 1, "/douban_tv_hot", summary="豆瓣热门电视剧", response_model=List[schemas.MediaInfo]
)
async def douban_tv_hot(
page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
热门电视剧 热门电视剧
""" """
@@ -133,7 +189,8 @@ async def douban_tv_hot(page: Optional[int] = 1,
@router.get("/tmdb_movies", summary="TMDB电影", response_model=List[schemas.MediaInfo]) @router.get("/tmdb_movies", summary="TMDB电影", response_model=List[schemas.MediaInfo])
async def tmdb_movies(sort_by: Optional[str] = "popularity.desc", async def tmdb_movies(
sort_by: Optional[str] = "popularity.desc",
with_genres: Optional[str] = "", with_genres: Optional[str] = "",
with_original_language: Optional[str] = "", with_original_language: Optional[str] = "",
with_keywords: Optional[str] = "", with_keywords: Optional[str] = "",
@@ -142,11 +199,13 @@ async def tmdb_movies(sort_by: Optional[str] = "popularity.desc",
vote_count: Optional[int] = 0, vote_count: Optional[int] = 0,
release_date: Optional[str] = "", release_date: Optional[str] = "",
page: Optional[int] = 1, page: Optional[int] = 1,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
浏览TMDB电影信息 浏览TMDB电影信息
""" """
return await RecommendChain().async_tmdb_movies(sort_by=sort_by, return await RecommendChain().async_tmdb_movies(
sort_by=sort_by,
with_genres=with_genres, with_genres=with_genres,
with_original_language=with_original_language, with_original_language=with_original_language,
with_keywords=with_keywords, with_keywords=with_keywords,
@@ -154,11 +213,13 @@ async def tmdb_movies(sort_by: Optional[str] = "popularity.desc",
vote_average=vote_average, vote_average=vote_average,
vote_count=vote_count, vote_count=vote_count,
release_date=release_date, release_date=release_date,
page=page) page=page,
)
@router.get("/tmdb_tvs", summary="TMDB剧集", response_model=List[schemas.MediaInfo]) @router.get("/tmdb_tvs", summary="TMDB剧集", response_model=List[schemas.MediaInfo])
async def tmdb_tvs(sort_by: Optional[str] = "popularity.desc", async def tmdb_tvs(
sort_by: Optional[str] = "popularity.desc",
with_genres: Optional[str] = "", with_genres: Optional[str] = "",
with_original_language: Optional[str] = "", with_original_language: Optional[str] = "",
with_keywords: Optional[str] = "", with_keywords: Optional[str] = "",
@@ -167,11 +228,13 @@ async def tmdb_tvs(sort_by: Optional[str] = "popularity.desc",
vote_count: Optional[int] = 0, vote_count: Optional[int] = 0,
release_date: Optional[str] = "", release_date: Optional[str] = "",
page: Optional[int] = 1, page: Optional[int] = 1,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
浏览TMDB剧集信息 浏览TMDB剧集信息
""" """
return await RecommendChain().async_tmdb_tvs(sort_by=sort_by, return await RecommendChain().async_tmdb_tvs(
sort_by=sort_by,
with_genres=with_genres, with_genres=with_genres,
with_original_language=with_original_language, with_original_language=with_original_language,
with_keywords=with_keywords, with_keywords=with_keywords,
@@ -179,12 +242,16 @@ async def tmdb_tvs(sort_by: Optional[str] = "popularity.desc",
vote_average=vote_average, vote_average=vote_average,
vote_count=vote_count, vote_count=vote_count,
release_date=release_date, release_date=release_date,
page=page) page=page,
)
@router.get("/tmdb_trending", summary="TMDB流行趋势", response_model=List[schemas.MediaInfo]) @router.get(
async def tmdb_trending(page: Optional[int] = 1, "/tmdb_trending", summary="TMDB流行趋势", response_model=List[schemas.MediaInfo]
_: schemas.TokenPayload = Depends(verify_token)) -> Any: )
async def tmdb_trending(
page: Optional[int] = 1, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
TMDB流行趋势 TMDB流行趋势
""" """
+268 -134
View File
@@ -47,17 +47,15 @@ def _merge_append_event(pending_event: Optional[dict], event: dict) -> dict:
return merged_event return merged_event
merged_event = dict(pending_event) merged_event = dict(pending_event)
merged_event.update({ merged_event.update({key: value for key, value in event.items() if key != "items"})
key: value
for key, value in event.items()
if key != "items"
})
merged_event["type"] = "append" merged_event["type"] = "append"
merged_event["items"] = [*(pending_event.get("items") or []), *items] merged_event["items"] = [*(pending_event.get("items") or []), *items]
return merged_event return merged_event
async def _iter_batched_search_events(event_source: AsyncIterator[dict]) -> AsyncIterator[dict]: async def _iter_batched_search_events(
event_source: AsyncIterator[dict],
) -> AsyncIterator[dict]:
""" """
对搜索流事件做轻量批处理,避免站点结果集中返回时产生过密 SSE。 对搜索流事件做轻量批处理,避免站点结果集中返回时产生过密 SSE。
""" """
@@ -90,7 +88,10 @@ async def _iter_batched_search_events(event_source: AsyncIterator[dict]) -> Asyn
if event.get("type") == "append": if event.get("type") == "append":
pending_append_event = _merge_append_event(pending_append_event, event) pending_append_event = _merge_append_event(pending_append_event, event)
if len(pending_append_event.get("items") or []) >= _SSE_APPEND_MAX_ITEMS: if (
len(pending_append_event.get("items") or [])
>= _SSE_APPEND_MAX_ITEMS
):
yield pending_append_event yield pending_append_event
pending_append_event = None pending_append_event = None
continue continue
@@ -121,20 +122,17 @@ async def _stream_search_events(request: Request, event_source: AsyncIterator[di
# 精确搜索会先发送 replace,再发送 done。done 再带整包 items 只会重复占用带宽和前端内存。 # 精确搜索会先发送 replace,再发送 done。done 再带整包 items 只会重复占用带宽和前端内存。
if event.get("type") == "replace" and event.get("items"): if event.get("type") == "replace" and event.get("items"):
has_sent_final_replace = True has_sent_final_replace = True
elif event.get("type") == "done" and has_sent_final_replace and event.get("stage") == "done" and event.get("items"): elif (
event = { event.get("type") == "done"
key: value and has_sent_final_replace
for key, value in event.items() and event.get("stage") == "done"
if key != "items" and event.get("items")
} ):
event = {key: value for key, value in event.items() if key != "items"}
yield _sse_event(event) yield _sse_event(event)
except Exception as err: except Exception as err:
logger.error(f"渐进式搜索出错:{err}", exc_info=True) logger.error(f"渐进式搜索出错:{err}", exc_info=True)
yield _sse_event({ yield _sse_event({"type": "error", "success": False, "message": str(err)})
"type": "error",
"success": False,
"message": str(err)
})
@router.get("/last", summary="查询搜索结果", response_model=List[schemas.Context]) @router.get("/last", summary="查询搜索结果", response_model=List[schemas.Context])
@@ -147,7 +145,8 @@ async def search_latest(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
@router.get("/media/{mediaid}/stream", summary="渐进式精确搜索资源") @router.get("/media/{mediaid}/stream", summary="渐进式精确搜索资源")
async def search_by_id_stream(request: Request, async def search_by_id_stream(
request: Request,
mediaid: str, mediaid: str,
mtype: Optional[str] = None, mtype: Optional[str] = None,
area: Optional[str] = "title", area: Optional[str] = "title",
@@ -155,7 +154,8 @@ async def search_by_id_stream(request: Request,
year: Optional[str] = None, year: Optional[str] = None,
season: Optional[str] = None, season: Optional[str] = None,
sites: Optional[str] = None, sites: Optional[str] = None,
_: schemas.TokenPayload = Depends(verify_resource_token)) -> Any: _: schemas.TokenPayload = Depends(verify_resource_token),
) -> Any:
""" """
根据TMDBID/豆瓣ID渐进式搜索站点资源,返回格式为SSE 根据TMDBID/豆瓣ID渐进式搜索站点资源,返回格式为SSE
""" """
@@ -172,77 +172,138 @@ async def search_by_id_stream(request: Request,
if mediaid.startswith("tmdb:"): if mediaid.startswith("tmdb:"):
tmdbid = int(mediaid.replace("tmdb:", "")) tmdbid = int(mediaid.replace("tmdb:", ""))
if settings.RECOGNIZE_SOURCE == "douban": if settings.RECOGNIZE_SOURCE == "douban":
doubaninfo = await media_chain.async_get_doubaninfo_by_tmdbid(tmdbid=tmdbid, mtype=media_type) doubaninfo = await media_chain.async_get_doubaninfo_by_tmdbid(
tmdbid=tmdbid, mtype=media_type
)
if doubaninfo: if doubaninfo:
torrents = search_chain.async_search_by_id_stream(doubanid=doubaninfo.get("id"), torrents = search_chain.async_search_by_id_stream(
mtype=media_type, area=area, doubanid=doubaninfo.get("id"),
season=media_season, sites=site_list, mtype=media_type,
cache_local=True) area=area,
season=media_season,
sites=site_list,
cache_local=True,
)
else: else:
yield {"type": "error", "success": False, "message": "未识别到豆瓣媒体信息"} yield {
"type": "error",
"success": False,
"message": "未识别到豆瓣媒体信息",
}
return return
else: else:
torrents = search_chain.async_search_by_id_stream(tmdbid=tmdbid, mtype=media_type, area=area, torrents = search_chain.async_search_by_id_stream(
season=media_season, sites=site_list, tmdbid=tmdbid,
cache_local=True) mtype=media_type,
area=area,
season=media_season,
sites=site_list,
cache_local=True,
)
elif mediaid.startswith("douban:"): elif mediaid.startswith("douban:"):
doubanid = mediaid.replace("douban:", "") doubanid = mediaid.replace("douban:", "")
if settings.RECOGNIZE_SOURCE == "themoviedb": if settings.RECOGNIZE_SOURCE == "themoviedb":
tmdbinfo = await media_chain.async_get_tmdbinfo_by_doubanid(doubanid=doubanid, mtype=media_type) tmdbinfo = await media_chain.async_get_tmdbinfo_by_doubanid(
doubanid=doubanid, mtype=media_type
)
if tmdbinfo: if tmdbinfo:
if tmdbinfo.get('season') and not media_season: if tmdbinfo.get("season") and not media_season:
media_season = tmdbinfo.get('season') media_season = tmdbinfo.get("season")
torrents = search_chain.async_search_by_id_stream(tmdbid=tmdbinfo.get("id"), torrents = search_chain.async_search_by_id_stream(
mtype=media_type, area=area, tmdbid=tmdbinfo.get("id"),
season=media_season, sites=site_list, mtype=media_type,
cache_local=True) area=area,
season=media_season,
sites=site_list,
cache_local=True,
)
else: else:
yield {"type": "error", "success": False, "message": "未识别到TMDB媒体信息"} yield {
"type": "error",
"success": False,
"message": "未识别到TMDB媒体信息",
}
return return
else: else:
torrents = search_chain.async_search_by_id_stream(doubanid=doubanid, mtype=media_type, area=area, torrents = search_chain.async_search_by_id_stream(
season=media_season, sites=site_list, doubanid=doubanid,
cache_local=True) mtype=media_type,
area=area,
season=media_season,
sites=site_list,
cache_local=True,
)
elif mediaid.startswith("bangumi:"): elif mediaid.startswith("bangumi:"):
bangumiid = int(mediaid.replace("bangumi:", "")) bangumiid = int(mediaid.replace("bangumi:", ""))
if settings.RECOGNIZE_SOURCE == "themoviedb": if settings.RECOGNIZE_SOURCE == "themoviedb":
tmdbinfo = await media_chain.async_get_tmdbinfo_by_bangumiid(bangumiid=bangumiid) tmdbinfo = await media_chain.async_get_tmdbinfo_by_bangumiid(
bangumiid=bangumiid
)
if tmdbinfo: if tmdbinfo:
torrents = search_chain.async_search_by_id_stream(tmdbid=tmdbinfo.get("id"), torrents = search_chain.async_search_by_id_stream(
mtype=media_type, area=area, tmdbid=tmdbinfo.get("id"),
season=media_season, sites=site_list, mtype=media_type,
cache_local=True) area=area,
season=media_season,
sites=site_list,
cache_local=True,
)
else: else:
yield {"type": "error", "success": False, "message": "未识别到TMDB媒体信息"} yield {
"type": "error",
"success": False,
"message": "未识别到TMDB媒体信息",
}
return return
else: else:
doubaninfo = await media_chain.async_get_doubaninfo_by_bangumiid(bangumiid=bangumiid) doubaninfo = await media_chain.async_get_doubaninfo_by_bangumiid(
bangumiid=bangumiid
)
if doubaninfo: if doubaninfo:
torrents = search_chain.async_search_by_id_stream(doubanid=doubaninfo.get("id"), torrents = search_chain.async_search_by_id_stream(
mtype=media_type, area=area, doubanid=doubaninfo.get("id"),
season=media_season, sites=site_list, mtype=media_type,
cache_local=True) area=area,
season=media_season,
sites=site_list,
cache_local=True,
)
else: else:
yield {"type": "error", "success": False, "message": "未识别到豆瓣媒体信息"} yield {
"type": "error",
"success": False,
"message": "未识别到豆瓣媒体信息",
}
return return
else: else:
event_data = MediaRecognizeConvertEventData( event_data = MediaRecognizeConvertEventData(
mediaid=mediaid, mediaid=mediaid, convert_type=settings.RECOGNIZE_SOURCE
convert_type=settings.RECOGNIZE_SOURCE )
event = await eventmanager.async_send_event(
ChainEventType.MediaRecognizeConvert, event_data
) )
event = await eventmanager.async_send_event(ChainEventType.MediaRecognizeConvert, event_data)
if event and event.event_data: if event and event.event_data:
event_data = event.event_data event_data = event.event_data
if event_data.media_dict: if event_data.media_dict:
search_id = event_data.media_dict.get("id") search_id = event_data.media_dict.get("id")
if event_data.convert_type == "themoviedb": if event_data.convert_type == "themoviedb":
torrents = search_chain.async_search_by_id_stream(tmdbid=search_id, mtype=media_type, torrents = search_chain.async_search_by_id_stream(
area=area, season=media_season, tmdbid=search_id,
sites=site_list, cache_local=True) mtype=media_type,
area=area,
season=media_season,
sites=site_list,
cache_local=True,
)
elif event_data.convert_type == "douban": elif event_data.convert_type == "douban":
torrents = search_chain.async_search_by_id_stream(doubanid=search_id, mtype=media_type, torrents = search_chain.async_search_by_id_stream(
area=area, season=media_season, doubanid=search_id,
sites=site_list, cache_local=True) mtype=media_type,
area=area,
season=media_season,
sites=site_list,
cache_local=True,
)
else: else:
if not title: if not title:
yield {"type": "error", "success": False, "message": "未知的媒体ID"} yield {"type": "error", "success": False, "message": "未知的媒体ID"}
@@ -261,15 +322,23 @@ async def search_by_id_stream(request: Request,
) )
if mediainfo: if mediainfo:
if settings.RECOGNIZE_SOURCE == "themoviedb": if settings.RECOGNIZE_SOURCE == "themoviedb":
torrents = search_chain.async_search_by_id_stream(tmdbid=mediainfo.tmdb_id, torrents = search_chain.async_search_by_id_stream(
mtype=media_type, area=area, tmdbid=mediainfo.tmdb_id,
season=media_season, sites=site_list, mtype=media_type,
cache_local=True) area=area,
season=media_season,
sites=site_list,
cache_local=True,
)
else: else:
torrents = search_chain.async_search_by_id_stream(doubanid=mediainfo.douban_id, torrents = search_chain.async_search_by_id_stream(
mtype=media_type, area=area, doubanid=mediainfo.douban_id,
season=media_season, sites=site_list, mtype=media_type,
cache_local=True) area=area,
season=media_season,
sites=site_list,
cache_local=True,
)
if not torrents: if not torrents:
yield {"type": "error", "success": False, "message": "未搜索到任何资源"} yield {"type": "error", "success": False, "message": "未搜索到任何资源"}
@@ -278,18 +347,22 @@ async def search_by_id_stream(request: Request,
async for event in torrents: async for event in torrents:
yield event yield event
return StreamingResponse(_stream_search_events(request, event_source()), media_type="text/event-stream") return StreamingResponse(
_stream_search_events(request, event_source()), media_type="text/event-stream"
)
@router.get("/media/{mediaid}", summary="精确搜索资源", response_model=schemas.Response) @router.get("/media/{mediaid}", summary="精确搜索资源", response_model=schemas.Response)
async def search_by_id(mediaid: str, async def search_by_id(
mediaid: str,
mtype: Optional[str] = None, mtype: Optional[str] = None,
area: Optional[str] = "title", area: Optional[str] = "title",
title: Optional[str] = None, title: Optional[str] = None,
year: Optional[str] = None, year: Optional[str] = None,
season: Optional[str] = None, season: Optional[str] = None,
sites: Optional[str] = None, sites: Optional[str] = None,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
根据TMDBID/豆瓣ID精确搜索站点资源 tmdb:/douban:/bangumi: 根据TMDBID/豆瓣ID精确搜索站点资源 tmdb:/douban:/bangumi:
""" """
@@ -313,72 +386,121 @@ async def search_by_id(mediaid: str,
tmdbid = int(mediaid.replace("tmdb:", "")) tmdbid = int(mediaid.replace("tmdb:", ""))
if settings.RECOGNIZE_SOURCE == "douban": if settings.RECOGNIZE_SOURCE == "douban":
# 通过TMDBID识别豆瓣ID # 通过TMDBID识别豆瓣ID
doubaninfo = await media_chain.async_get_doubaninfo_by_tmdbid(tmdbid=tmdbid, mtype=media_type) doubaninfo = await media_chain.async_get_doubaninfo_by_tmdbid(
tmdbid=tmdbid, mtype=media_type
)
if doubaninfo: if doubaninfo:
torrents = await search_chain.async_search_by_id(doubanid=doubaninfo.get("id"), torrents = await search_chain.async_search_by_id(
mtype=media_type, area=area, season=media_season, doubanid=doubaninfo.get("id"),
sites=site_list, cache_local=True) mtype=media_type,
area=area,
season=media_season,
sites=site_list,
cache_local=True,
)
else: else:
return schemas.Response(success=False, message="未识别到豆瓣媒体信息") return schemas.Response(success=False, message="未识别到豆瓣媒体信息")
else: else:
torrents = await search_chain.async_search_by_id(tmdbid=tmdbid, mtype=media_type, area=area, torrents = await search_chain.async_search_by_id(
tmdbid=tmdbid,
mtype=media_type,
area=area,
season=media_season, season=media_season,
sites=site_list, cache_local=True) sites=site_list,
cache_local=True,
)
elif mediaid.startswith("douban:"): elif mediaid.startswith("douban:"):
doubanid = mediaid.replace("douban:", "") doubanid = mediaid.replace("douban:", "")
if settings.RECOGNIZE_SOURCE == "themoviedb": if settings.RECOGNIZE_SOURCE == "themoviedb":
# 通过豆瓣ID识别TMDBID # 通过豆瓣ID识别TMDBID
tmdbinfo = await media_chain.async_get_tmdbinfo_by_doubanid(doubanid=doubanid, mtype=media_type) tmdbinfo = await media_chain.async_get_tmdbinfo_by_doubanid(
doubanid=doubanid, mtype=media_type
)
if tmdbinfo: if tmdbinfo:
if tmdbinfo.get('season') and not media_season: if tmdbinfo.get("season") and not media_season:
media_season = tmdbinfo.get('season') media_season = tmdbinfo.get("season")
torrents = await search_chain.async_search_by_id(tmdbid=tmdbinfo.get("id"), torrents = await search_chain.async_search_by_id(
mtype=media_type, area=area, season=media_season, tmdbid=tmdbinfo.get("id"),
sites=site_list, cache_local=True) mtype=media_type,
area=area,
season=media_season,
sites=site_list,
cache_local=True,
)
else: else:
return schemas.Response(success=False, message="未识别到TMDB媒体信息") return schemas.Response(success=False, message="未识别到TMDB媒体信息")
else: else:
torrents = await search_chain.async_search_by_id(doubanid=doubanid, mtype=media_type, area=area, torrents = await search_chain.async_search_by_id(
doubanid=doubanid,
mtype=media_type,
area=area,
season=media_season, season=media_season,
sites=site_list, cache_local=True) sites=site_list,
cache_local=True,
)
elif mediaid.startswith("bangumi:"): elif mediaid.startswith("bangumi:"):
bangumiid = int(mediaid.replace("bangumi:", "")) bangumiid = int(mediaid.replace("bangumi:", ""))
if settings.RECOGNIZE_SOURCE == "themoviedb": if settings.RECOGNIZE_SOURCE == "themoviedb":
# 通过BangumiID识别TMDBID # 通过BangumiID识别TMDBID
tmdbinfo = await media_chain.async_get_tmdbinfo_by_bangumiid(bangumiid=bangumiid) tmdbinfo = await media_chain.async_get_tmdbinfo_by_bangumiid(
bangumiid=bangumiid
)
if tmdbinfo: if tmdbinfo:
torrents = await search_chain.async_search_by_id(tmdbid=tmdbinfo.get("id"), torrents = await search_chain.async_search_by_id(
mtype=media_type, area=area, season=media_season, tmdbid=tmdbinfo.get("id"),
sites=site_list, cache_local=True) mtype=media_type,
area=area,
season=media_season,
sites=site_list,
cache_local=True,
)
else: else:
return schemas.Response(success=False, message="未识别到TMDB媒体信息") return schemas.Response(success=False, message="未识别到TMDB媒体信息")
else: else:
# 通过BangumiID识别豆瓣ID # 通过BangumiID识别豆瓣ID
doubaninfo = await media_chain.async_get_doubaninfo_by_bangumiid(bangumiid=bangumiid) doubaninfo = await media_chain.async_get_doubaninfo_by_bangumiid(
bangumiid=bangumiid
)
if doubaninfo: if doubaninfo:
torrents = await search_chain.async_search_by_id(doubanid=doubaninfo.get("id"), torrents = await search_chain.async_search_by_id(
mtype=media_type, area=area, season=media_season, doubanid=doubaninfo.get("id"),
sites=site_list, cache_local=True) mtype=media_type,
area=area,
season=media_season,
sites=site_list,
cache_local=True,
)
else: else:
return schemas.Response(success=False, message="未识别到豆瓣媒体信息") return schemas.Response(success=False, message="未识别到豆瓣媒体信息")
else: else:
# 未知前缀,广播事件解析媒体信息 # 未知前缀,广播事件解析媒体信息
event_data = MediaRecognizeConvertEventData( event_data = MediaRecognizeConvertEventData(
mediaid=mediaid, mediaid=mediaid, convert_type=settings.RECOGNIZE_SOURCE
convert_type=settings.RECOGNIZE_SOURCE )
event = await eventmanager.async_send_event(
ChainEventType.MediaRecognizeConvert, event_data
) )
event = await eventmanager.async_send_event(ChainEventType.MediaRecognizeConvert, event_data)
# 使用事件返回的上下文数据 # 使用事件返回的上下文数据
if event and event.event_data: if event and event.event_data:
event_data: MediaRecognizeConvertEventData = event.event_data event_data: MediaRecognizeConvertEventData = event.event_data
if event_data.media_dict: if event_data.media_dict:
search_id = event_data.media_dict.get("id") search_id = event_data.media_dict.get("id")
if event_data.convert_type == "themoviedb": if event_data.convert_type == "themoviedb":
torrents = await search_chain.async_search_by_id(tmdbid=search_id, mtype=media_type, area=area, torrents = await search_chain.async_search_by_id(
season=media_season, cache_local=True) tmdbid=search_id,
mtype=media_type,
area=area,
season=media_season,
cache_local=True,
)
elif event_data.convert_type == "douban": elif event_data.convert_type == "douban":
torrents = await search_chain.async_search_by_id(doubanid=search_id, mtype=media_type, area=area, torrents = await search_chain.async_search_by_id(
season=media_season, cache_local=True) doubanid=search_id,
mtype=media_type,
area=area,
season=media_season,
cache_local=True,
)
else: else:
if not title: if not title:
return schemas.Response(success=False, message="未知的媒体ID") return schemas.Response(success=False, message="未知的媒体ID")
@@ -397,63 +519,79 @@ async def search_by_id(mediaid: str,
) )
if mediainfo: if mediainfo:
if settings.RECOGNIZE_SOURCE == "themoviedb": if settings.RECOGNIZE_SOURCE == "themoviedb":
torrents = await search_chain.async_search_by_id(tmdbid=mediainfo.tmdb_id, mtype=media_type, torrents = await search_chain.async_search_by_id(
tmdbid=mediainfo.tmdb_id,
mtype=media_type,
area=area, area=area,
season=media_season, cache_local=True) season=media_season,
cache_local=True,
)
else: else:
torrents = await search_chain.async_search_by_id(doubanid=mediainfo.douban_id, mtype=media_type, torrents = await search_chain.async_search_by_id(
doubanid=mediainfo.douban_id,
mtype=media_type,
area=area, area=area,
season=media_season, cache_local=True) season=media_season,
cache_local=True,
)
# 返回搜索结果 # 返回搜索结果
if not torrents: if not torrents:
return schemas.Response(success=False, message="未搜索到任何资源") return schemas.Response(success=False, message="未搜索到任何资源")
else: else:
return schemas.Response(success=True, data=[torrent.to_dict() for torrent in torrents]) return schemas.Response(
success=True, data=[torrent.to_dict() for torrent in torrents]
)
@router.get("/title/stream", summary="渐进式模糊搜索资源") @router.get("/title/stream", summary="渐进式模糊搜索资源")
async def search_by_title_stream(request: Request, async def search_by_title_stream(
request: Request,
keyword: Optional[str] = None, keyword: Optional[str] = None,
page: Optional[int] = 0, page: Optional[int] = 0,
sites: Optional[str] = None, sites: Optional[str] = None,
_: schemas.TokenPayload = Depends(verify_resource_token)) -> Any: _: schemas.TokenPayload = Depends(verify_resource_token),
) -> Any:
""" """
根据名称渐进式模糊搜索站点资源,返回格式为SSE 根据名称渐进式模糊搜索站点资源,返回格式为SSE
""" """
event_source = SearchChain().async_search_by_title_stream( event_source = SearchChain().async_search_by_title_stream(
title=keyword, title=keyword, page=page, sites=_parse_site_list(sites), cache_local=True
page=page, )
sites=_parse_site_list(sites), return StreamingResponse(
cache_local=True _stream_search_events(request, event_source), media_type="text/event-stream"
) )
return StreamingResponse(_stream_search_events(request, event_source), media_type="text/event-stream")
@router.get("/title", summary="模糊搜索资源", response_model=schemas.Response) @router.get("/title", summary="模糊搜索资源", response_model=schemas.Response)
async def search_by_title(keyword: Optional[str] = None, async def search_by_title(
keyword: Optional[str] = None,
page: Optional[int] = 0, page: Optional[int] = 0,
sites: Optional[str] = None, sites: Optional[str] = None,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
根据名称模糊搜索站点资源,支持分页,关键词为空是返回首页资源 根据名称模糊搜索站点资源,支持分页,关键词为空是返回首页资源
""" """
torrents = await SearchChain().async_search_by_title( torrents = await SearchChain().async_search_by_title(
title=keyword, page=page, title=keyword, page=page, sites=_parse_site_list(sites), cache_local=True
sites=_parse_site_list(sites),
cache_local=True
) )
if not torrents: if not torrents:
return schemas.Response(success=False, message="未搜索到任何资源") return schemas.Response(success=False, message="未搜索到任何资源")
return schemas.Response(success=True, data=[torrent.to_dict() for torrent in torrents]) return schemas.Response(
success=True, data=[torrent.to_dict() for torrent in torrents]
)
@router.post("/recommend", summary="AI推荐资源", response_model=schemas.Response) @router.post("/recommend", summary="AI推荐资源", response_model=schemas.Response)
async def recommend_search_results( async def recommend_search_results(
filtered_indices: Optional[List[int]] = Body(None, embed=True, description="筛选后的索引列表"), filtered_indices: Optional[List[int]] = Body(
None, embed=True, description="筛选后的索引列表"
),
check_only: bool = Body(False, embed=True, description="仅检查状态,不启动新任务"), check_only: bool = Body(False, embed=True, description="仅检查状态,不启动新任务"),
force: bool = Body(False, embed=True, description="强制重新推荐,清除旧结果"), force: bool = Body(False, embed=True, description="强制重新推荐,清除旧结果"),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
AI推荐资源 - 轮询接口 AI推荐资源 - 轮询接口
前端轮询此接口,发送筛选后的索引(如果有筛选) 前端轮询此接口,发送筛选后的索引(如果有筛选)
@@ -477,9 +615,9 @@ async def recommend_search_results(
# 从缓存获取上次搜索结果 # 从缓存获取上次搜索结果
results = await SearchChain().async_last_search_results() or [] results = await SearchChain().async_last_search_results() or []
if not results: if not results:
return schemas.Response(success=False, message="没有可用的搜索结果", data={ return schemas.Response(
"status": "error" success=False, message="没有可用的搜索结果", data={"status": "error"}
}) )
recommend_chain = SearchChain() recommend_chain = SearchChain()
@@ -487,16 +625,12 @@ async def recommend_search_results(
if force: if force:
# 检查功能是否启用 # 检查功能是否启用
if not recommend_chain.is_ai_recommend_enabled: if not recommend_chain.is_ai_recommend_enabled:
return schemas.Response(success=True, data={ return schemas.Response(success=True, data={"status": "disabled"})
"status": "disabled"
})
logger.info("收到新推荐请求,清除旧结果并启动新任务") logger.info("收到新推荐请求,清除旧结果并启动新任务")
recommend_chain.cancel_ai_recommend() recommend_chain.cancel_ai_recommend()
recommend_chain.start_recommend_task(filtered_indices, len(results), results) recommend_chain.start_recommend_task(filtered_indices, len(results), results)
# 直接返回运行中状态 # 直接返回运行中状态
return schemas.Response(success=True, data={ return schemas.Response(success=True, data={"status": "running"})
"status": "running"
})
# 如果是仅检查模式,不传递 filtered_indices(避免触发请求变化检测) # 如果是仅检查模式,不传递 filtered_indices(避免触发请求变化检测)
if check_only: if check_only:
@@ -505,7 +639,9 @@ async def recommend_search_results(
# 如果有错误,将错误信息放到message中 # 如果有错误,将错误信息放到message中
if current_status.get("status") == "error": if current_status.get("status") == "error":
error_msg = current_status.pop("error", "未知错误") error_msg = current_status.pop("error", "未知错误")
return schemas.Response(success=False, message=error_msg, data=current_status) return schemas.Response(
success=False, message=error_msg, data=current_status
)
return schemas.Response(success=True, data=current_status) return schemas.Response(success=True, data=current_status)
# 获取当前状态(会检测请求是否变化) # 获取当前状态(会检测请求是否变化)
@@ -519,9 +655,7 @@ async def recommend_search_results(
if status_data["status"] == "idle": if status_data["status"] == "idle":
recommend_chain.start_recommend_task(filtered_indices, len(results), results) recommend_chain.start_recommend_task(filtered_indices, len(results), results)
# 立即返回运行中状态 # 立即返回运行中状态
return schemas.Response(success=True, data={ return schemas.Response(success=True, data={"status": "running"})
"status": "running"
})
# 如果有错误,将错误信息放到message中 # 如果有错误,将错误信息放到message中
if status_data.get("status") == "error": if status_data.get("status") == "error":
+118 -68
View File
@@ -21,7 +21,10 @@ from app.db.models.sitestatistic import SiteStatistic
from app.db.models.siteuserdata import SiteUserData from app.db.models.siteuserdata import SiteUserData
from app.db.site_oper import SiteOper from app.db.site_oper import SiteOper
from app.db.systemconfig_oper import SystemConfigOper from app.db.systemconfig_oper import SystemConfigOper
from app.db.user_oper import get_current_active_superuser, get_current_active_superuser_async from app.db.user_oper import (
get_current_active_superuser,
get_current_active_superuser_async,
)
from app.helper.sites import SitesHelper # noqa from app.helper.sites import SitesHelper # noqa
from app.scheduler import Scheduler from app.scheduler import Scheduler
from app.schemas.types import SystemConfigKey, EventType from app.schemas.types import SystemConfigKey, EventType
@@ -31,8 +34,10 @@ router = APIRouter()
@router.get("/", summary="所有站点", response_model=List[schemas.Site]) @router.get("/", summary="所有站点", response_model=List[schemas.Site])
async def read_sites(db: AsyncSession = Depends(get_async_db), async def read_sites(
_: User = Depends(get_current_active_superuser)) -> List[dict]: db: AsyncSession = Depends(get_async_db),
_: User = Depends(get_current_active_superuser),
) -> List[dict]:
""" """
获取站点列表 获取站点列表
""" """
@@ -44,7 +49,7 @@ async def add_site(
*, *,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
site_in: schemas.Site, site_in: schemas.Site,
_: User = Depends(get_current_active_superuser) _: User = Depends(get_current_active_superuser),
) -> Any: ) -> Any:
""" """
新增站点 新增站点
@@ -52,11 +57,15 @@ async def add_site(
if not site_in.url: if not site_in.url:
return schemas.Response(success=False, message="站点地址不能为空") return schemas.Response(success=False, message="站点地址不能为空")
if SitesHelper().auth_level < 2: if SitesHelper().auth_level < 2:
return schemas.Response(success=False, message="用户未通过认证,无法使用站点功能!") return schemas.Response(
success=False, message="用户未通过认证,无法使用站点功能!"
)
domain = StringUtils.get_url_domain(site_in.url) domain = StringUtils.get_url_domain(site_in.url)
site_info = await SitesHelper().async_get_indexer(domain) site_info = await SitesHelper().async_get_indexer(domain)
if not site_info: if not site_info:
return schemas.Response(success=False, message="该站点不支持,请检查站点域名是否正确") return schemas.Response(
success=False, message="该站点不支持,请检查站点域名是否正确"
)
if await Site.async_get_by_domain(db, domain): if await Site.async_get_by_domain(db, domain):
return schemas.Response(success=False, message=f"{domain} 站点己存在") return schemas.Response(success=False, message=f"{domain} 站点己存在")
# 保存站点信息 # 保存站点信息
@@ -70,9 +79,7 @@ async def add_site(
site = Site(**site_in.model_dump()) site = Site(**site_in.model_dump())
site.create(db) site.create(db)
# 通知站点更新 # 通知站点更新
await eventmanager.async_send_event(EventType.SiteUpdated, { await eventmanager.async_send_event(EventType.SiteUpdated, {"domain": domain})
"domain": domain
})
return schemas.Response(success=True) return schemas.Response(success=True)
@@ -81,7 +88,7 @@ async def update_site(
*, *,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
site_in: schemas.Site, site_in: schemas.Site,
_: User = Depends(get_current_active_superuser) _: User = Depends(get_current_active_superuser),
) -> Any: ) -> Any:
""" """
更新站点信息 更新站点信息
@@ -95,18 +102,23 @@ async def update_site(
site_in.domain = StringUtils.get_url_domain(site_in.url) site_in.domain = StringUtils.get_url_domain(site_in.url)
await site.async_update(db, site_in.model_dump()) await site.async_update(db, site_in.model_dump())
# 通知站点更新 # 通知站点更新
await eventmanager.async_send_event(EventType.SiteUpdated, { await eventmanager.async_send_event(
EventType.SiteUpdated,
{
"site_id": site_in.id, "site_id": site_in.id,
"domain": site_in.domain, "domain": site_in.domain,
"name": site_in.name, "name": site_in.name,
"site_url": site_in.url "site_url": site_in.url,
}) },
)
return schemas.Response(success=True) return schemas.Response(success=True)
@router.get("/cookiecloud", summary="CookieCloud同步", response_model=schemas.Response) @router.get("/cookiecloud", summary="CookieCloud同步", response_model=schemas.Response)
async def cookie_cloud_sync(background_tasks: BackgroundTasks, async def cookie_cloud_sync(
_: User = Depends(get_current_active_superuser_async)) -> Any: background_tasks: BackgroundTasks,
_: User = Depends(get_current_active_superuser_async),
) -> Any:
""" """
运行CookieCloud同步站点信息 运行CookieCloud同步站点信息
""" """
@@ -115,8 +127,9 @@ async def cookie_cloud_sync(background_tasks: BackgroundTasks,
@router.get("/reset", summary="重置站点", response_model=schemas.Response) @router.get("/reset", summary="重置站点", response_model=schemas.Response)
def reset(db: AsyncSession = Depends(get_db), def reset(
_: User = Depends(get_current_active_superuser)) -> Any: db: AsyncSession = Depends(get_db), _: User = Depends(get_current_active_superuser)
) -> Any:
""" """
清空所有站点数据并重新同步CookieCloud站点信息 清空所有站点数据并重新同步CookieCloud站点信息
""" """
@@ -126,18 +139,18 @@ def reset(db: AsyncSession = Depends(get_db),
# 启动定时服务 # 启动定时服务
Scheduler().start("cookiecloud", manual=True) Scheduler().start("cookiecloud", manual=True)
# 插件站点删除 # 插件站点删除
eventmanager.send_event(EventType.SiteDeleted, eventmanager.send_event(EventType.SiteDeleted, {"site_id": "*"})
{
"site_id": "*"
})
return schemas.Response(success=True, message="站点已重置!") return schemas.Response(success=True, message="站点已重置!")
@router.post("/priorities", summary="批量更新站点优先级", response_model=schemas.Response) @router.post(
"/priorities", summary="批量更新站点优先级", response_model=schemas.Response
)
async def update_sites_priority( async def update_sites_priority(
priorities: List[dict], priorities: List[dict],
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: User = Depends(get_current_active_superuser_async)) -> Any: _: User = Depends(get_current_active_superuser_async),
) -> Any:
""" """
批量更新站点优先级 批量更新站点优先级
""" """
@@ -148,14 +161,17 @@ async def update_sites_priority(
return schemas.Response(success=True) return schemas.Response(success=True)
@router.get("/cookie/{site_id}", summary="更新站点Cookie&UA", response_model=schemas.Response) @router.get(
"/cookie/{site_id}", summary="更新站点Cookie&UA", response_model=schemas.Response
)
def update_cookie( def update_cookie(
site_id: int, site_id: int,
username: str, username: str,
password: str, password: str,
code: Optional[str] = None, code: Optional[str] = None,
db: Session = Depends(get_db), db: Session = Depends(get_db),
_: User = Depends(get_current_active_superuser)) -> Any: _: User = Depends(get_current_active_superuser),
) -> Any:
""" """
使用用户密码更新站点Cookie 使用用户密码更新站点Cookie
""" """
@@ -167,18 +183,20 @@ def update_cookie(
detail=f"站点 {site_id} 不存在!", detail=f"站点 {site_id} 不存在!",
) )
# 更新Cookie # 更新Cookie
state, message = SiteChain().update_cookie(site_info=site_info, state, message = SiteChain().update_cookie(
username=username, site_info=site_info, username=username, password=password, two_step_code=code
password=password, )
two_step_code=code)
return schemas.Response(success=state, message=message) return schemas.Response(success=state, message=message)
@router.post("/userdata/{site_id}", summary="更新站点用户数据", response_model=schemas.Response) @router.post(
"/userdata/{site_id}", summary="更新站点用户数据", response_model=schemas.Response
)
def refresh_userdata( def refresh_userdata(
site_id: int, site_id: int,
db: Session = Depends(get_db), db: Session = Depends(get_db),
_: User = Depends(get_current_active_superuser)) -> Any: _: User = Depends(get_current_active_superuser),
) -> Any:
""" """
刷新站点用户数据 刷新站点用户数据
""" """
@@ -190,15 +208,22 @@ def refresh_userdata(
) )
indexer = SitesHelper().get_indexer(site.domain) indexer = SitesHelper().get_indexer(site.domain)
if not indexer: if not indexer:
return schemas.Response(success=False, message="站点不支持索引或未通过用户认证!") return schemas.Response(
success=False, message="站点不支持索引或未通过用户认证!"
)
user_data = SiteChain().refresh_userdata(site=indexer) or {} user_data = SiteChain().refresh_userdata(site=indexer) or {}
return schemas.Response(success=True, data=user_data) return schemas.Response(success=True, data=user_data)
@router.get("/userdata/latest", summary="查询所有站点最新用户数据", response_model=List[schemas.SiteUserData]) @router.get(
"/userdata/latest",
summary="查询所有站点最新用户数据",
response_model=List[schemas.SiteUserData],
)
async def read_userdata_latest( async def read_userdata_latest(
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: User = Depends(get_current_active_superuser_async)) -> Any: _: User = Depends(get_current_active_superuser_async),
) -> Any:
""" """
查询所有站点最新用户数据 查询所有站点最新用户数据
""" """
@@ -208,12 +233,15 @@ async def read_userdata_latest(
return [user_data.to_dict() for user_data in user_datas] return [user_data.to_dict() for user_data in user_datas]
@router.get("/userdata/{site_id}", summary="查询某站点用户数据", response_model=schemas.Response) @router.get(
"/userdata/{site_id}", summary="查询某站点用户数据", response_model=schemas.Response
)
async def read_userdata( async def read_userdata(
site_id: int, site_id: int,
workdate: Optional[str] = None, workdate: Optional[str] = None,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: User = Depends(get_current_active_superuser_async)) -> Any: _: User = Depends(get_current_active_superuser_async),
) -> Any:
""" """
查询站点用户数据 查询站点用户数据
""" """
@@ -223,16 +251,20 @@ async def read_userdata(
status_code=404, status_code=404,
detail=f"站点 {site_id} 不存在", detail=f"站点 {site_id} 不存在",
) )
user_datas = await SiteUserData.async_get_by_domain(db, domain=site.domain, workdate=workdate) user_datas = await SiteUserData.async_get_by_domain(
db, domain=site.domain, workdate=workdate
)
if not user_datas: if not user_datas:
return schemas.Response(success=False, data=[]) return schemas.Response(success=False, data=[])
return schemas.Response(success=True, data=[data.to_dict() for data in user_datas]) return schemas.Response(success=True, data=[data.to_dict() for data in user_datas])
@router.get("/test/{site_id}", summary="连接测试", response_model=schemas.Response) @router.get("/test/{site_id}", summary="连接测试", response_model=schemas.Response)
def test_site(site_id: int, def test_site(
site_id: int,
db: Session = Depends(get_db), db: Session = Depends(get_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
测试站点是否可用 测试站点是否可用
""" """
@@ -247,9 +279,11 @@ def test_site(site_id: int,
@router.get("/icon/{site_id}", summary="站点图标", response_model=schemas.Response) @router.get("/icon/{site_id}", summary="站点图标", response_model=schemas.Response)
async def site_icon(site_id: int, async def site_icon(
site_id: int,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
获取站点图标:base64或者url 获取站点图标:base64或者url
""" """
@@ -262,15 +296,19 @@ async def site_icon(site_id: int,
icon = await SiteIcon.async_get_by_domain(db, site.domain) icon = await SiteIcon.async_get_by_domain(db, site.domain)
if not icon: if not icon:
return schemas.Response(success=False, message="站点图标不存在!") return schemas.Response(success=False, message="站点图标不存在!")
return schemas.Response(success=True, data={ return schemas.Response(
"icon": icon.base64 if icon.base64 else icon.url success=True, data={"icon": icon.base64 if icon.base64 else icon.url}
}) )
@router.get("/category/{site_id}", summary="站点分类", response_model=List[schemas.SiteCategory]) @router.get(
async def site_category(site_id: int, "/category/{site_id}", summary="站点分类", response_model=List[schemas.SiteCategory]
)
async def site_category(
site_id: int,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
获取站点分类 获取站点分类
""" """
@@ -286,7 +324,7 @@ async def site_category(site_id: int,
status_code=404, status_code=404,
detail=f"站点 {site.domain} 不支持", detail=f"站点 {site.domain} 不支持",
) )
category: Dict[str, List[dict]] = indexer.get('category') or [] category: Dict[str, List[dict]] = indexer.get("category") or []
if not category: if not category:
return [] return []
result = [] result = []
@@ -297,13 +335,17 @@ async def site_category(site_id: int,
return result return result
@router.get("/resource/{site_id}", summary="站点资源", response_model=List[schemas.TorrentInfo]) @router.get(
async def site_resource(site_id: int, "/resource/{site_id}", summary="站点资源", response_model=List[schemas.TorrentInfo]
)
async def site_resource(
site_id: int,
keyword: Optional[str] = None, keyword: Optional[str] = None,
cat: Optional[str] = None, cat: Optional[str] = None,
page: Optional[int] = 0, page: Optional[int] = 0,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: User = Depends(get_current_active_superuser_async)) -> Any: _: User = Depends(get_current_active_superuser_async),
) -> Any:
""" """
浏览站点资源 浏览站点资源
""" """
@@ -313,7 +355,9 @@ async def site_resource(site_id: int,
status_code=404, status_code=404,
detail=f"站点 {site_id} 不存在", detail=f"站点 {site_id} 不存在",
) )
torrents = await TorrentsChain().async_browse(domain=site.domain, keyword=keyword, cat=cat, page=page) torrents = await TorrentsChain().async_browse(
domain=site.domain, keyword=keyword, cat=cat, page=page
)
if not torrents: if not torrents:
return [] return []
return [torrent.to_dict() for torrent in torrents] return [torrent.to_dict() for torrent in torrents]
@@ -323,7 +367,7 @@ async def site_resource(site_id: int,
async def read_site_by_domain( async def read_site_by_domain(
site_url: str, site_url: str,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token) _: schemas.TokenPayload = Depends(verify_token),
) -> Any: ) -> Any:
""" """
通过域名获取站点信息 通过域名获取站点信息
@@ -338,11 +382,15 @@ async def read_site_by_domain(
return site return site
@router.get("/statistic/{site_url}", summary="特定站点统计信息", response_model=schemas.SiteStatistic) @router.get(
"/statistic/{site_url}",
summary="特定站点统计信息",
response_model=schemas.SiteStatistic,
)
async def read_statistic_by_domain( async def read_statistic_by_domain(
site_url: str, site_url: str,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token) _: schemas.TokenPayload = Depends(verify_token),
) -> Any: ) -> Any:
""" """
通过域名获取站点统计信息 通过域名获取站点统计信息
@@ -354,10 +402,12 @@ async def read_statistic_by_domain(
return schemas.SiteStatistic(domain=domain) return schemas.SiteStatistic(domain=domain)
@router.get("/statistic", summary="所有站点统计信息", response_model=List[schemas.SiteStatistic]) @router.get(
"/statistic", summary="所有站点统计信息", response_model=List[schemas.SiteStatistic]
)
async def read_statistics( async def read_statistics(
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token) _: schemas.TokenPayload = Depends(verify_token),
) -> Any: ) -> Any:
""" """
获取所有站点统计信息 获取所有站点统计信息
@@ -366,8 +416,10 @@ async def read_statistics(
@router.get("/rss", summary="所有订阅站点", response_model=List[schemas.Site]) @router.get("/rss", summary="所有订阅站点", response_model=List[schemas.Site])
async def read_rss_sites(db: AsyncSession = Depends(get_async_db), async def read_rss_sites(
_: schemas.TokenPayload = Depends(verify_token)) -> List[dict]: db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token),
) -> List[dict]:
""" """
获取站点列表 获取站点列表
""" """
@@ -394,8 +446,7 @@ async def read_auth_sites(_: schemas.TokenPayload = Depends(verify_token)) -> di
@router.post("/auth", summary="用户站点认证", response_model=schemas.Response) @router.post("/auth", summary="用户站点认证", response_model=schemas.Response)
def auth_site( def auth_site(
auth_info: schemas.SiteAuth, auth_info: schemas.SiteAuth, _: User = Depends(get_current_active_superuser)
_: User = Depends(get_current_active_superuser)
) -> Any: ) -> Any:
""" """
用户站点认证 用户站点认证
@@ -412,7 +463,9 @@ def auth_site(
return schemas.Response(success=status, message=msg) return schemas.Response(success=status, message=msg)
@router.get("/mapping", summary="获取站点域名到名称的映射", response_model=schemas.Response) @router.get(
"/mapping", summary="获取站点域名到名称的映射", response_model=schemas.Response
)
async def site_mapping(_: User = Depends(get_current_active_superuser_async)): async def site_mapping(_: User = Depends(get_current_active_superuser_async)):
""" """
获取站点域名到名称的映射关系 获取站点域名到名称的映射关系
@@ -439,7 +492,7 @@ async def support_sites(_: User = Depends(get_current_active_superuser_async)):
async def read_site( async def read_site(
site_id: int, site_id: int,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: User = Depends(get_current_active_superuser_async) _: User = Depends(get_current_active_superuser_async),
) -> Any: ) -> Any:
""" """
通过ID获取站点信息 通过ID获取站点信息
@@ -457,15 +510,12 @@ async def read_site(
async def delete_site( async def delete_site(
site_id: int, site_id: int,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: User = Depends(get_current_active_superuser_async) _: User = Depends(get_current_active_superuser_async),
) -> Any: ) -> Any:
""" """
删除站点 删除站点
""" """
await Site.async_delete(db, site_id) await Site.async_delete(db, site_id)
# 插件站点删除 # 插件站点删除
await eventmanager.async_send_event(EventType.SiteDeleted, await eventmanager.async_send_event(EventType.SiteDeleted, {"site_id": site_id})
{
"site_id": site_id
})
return schemas.Response(success=True) return schemas.Response(success=True)
+64 -35
View File
@@ -12,7 +12,10 @@ from app.chain.transfer import TransferChain
from app.core.config import settings from app.core.config import settings
from app.core.security import verify_token from app.core.security import verify_token
from app.db.models import User from app.db.models import User
from app.db.user_oper import get_current_active_superuser, get_current_active_superuser_async from app.db.user_oper import (
get_current_active_superuser,
get_current_active_superuser_async,
)
from app.helper.progress import ProgressHelper from app.helper.progress import ProgressHelper
from app.schemas.types import ProgressKey from app.schemas.types import ProgressKey
from app.utils.string import StringUtils from app.utils.string import StringUtils
@@ -31,7 +34,9 @@ def qrcode(name: str, _: schemas.TokenPayload = Depends(verify_token)) -> Any:
return schemas.Response(success=False, message=errmsg) return schemas.Response(success=False, message=errmsg)
@router.get("/auth_url/{name}", summary="获取 OAuth2 授权 URL", response_model=schemas.Response) @router.get(
"/auth_url/{name}", summary="获取 OAuth2 授权 URL", response_model=schemas.Response
)
def auth_url(name: str, _: schemas.TokenPayload = Depends(verify_token)) -> Any: def auth_url(name: str, _: schemas.TokenPayload = Depends(verify_token)) -> Any:
""" """
获取 OAuth2 授权 URL 获取 OAuth2 授权 URL
@@ -43,8 +48,12 @@ def auth_url(name: str, _: schemas.TokenPayload = Depends(verify_token)) -> Any:
@router.get("/check/{name}", summary="二维码登录确认", response_model=schemas.Response) @router.get("/check/{name}", summary="二维码登录确认", response_model=schemas.Response)
def check(name: str, ck: Optional[str] = None, t: Optional[str] = None, def check(
_: schemas.TokenPayload = Depends(verify_token)) -> Any: name: str,
ck: Optional[str] = None,
t: Optional[str] = None,
_: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
二维码登录确认 二维码登录确认
""" """
@@ -58,9 +67,7 @@ def check(name: str, ck: Optional[str] = None, t: Optional[str] = None,
@router.post("/save/{name}", summary="保存存储配置", response_model=schemas.Response) @router.post("/save/{name}", summary="保存存储配置", response_model=schemas.Response)
def save(name: str, def save(name: str, conf: dict, _: User = Depends(get_current_active_superuser)) -> Any:
conf: dict,
_: User = Depends(get_current_active_superuser)) -> Any:
""" """
保存存储配置 保存存储配置
""" """
@@ -69,8 +76,7 @@ def save(name: str,
@router.get("/reset/{name}", summary="重置存储配置", response_model=schemas.Response) @router.get("/reset/{name}", summary="重置存储配置", response_model=schemas.Response)
def reset(name: str, def reset(name: str, _: User = Depends(get_current_active_superuser)) -> Any:
_: User = Depends(get_current_active_superuser)) -> Any:
""" """
重置存储配置 重置存储配置
""" """
@@ -79,9 +85,11 @@ def reset(name: str,
@router.post("/list", summary="所有目录和文件", response_model=List[schemas.FileItem]) @router.post("/list", summary="所有目录和文件", response_model=List[schemas.FileItem])
def list_files(fileitem: schemas.FileItem, def list_files(
sort: Optional[str] = 'updated_at', fileitem: schemas.FileItem,
_: User = Depends(get_current_active_superuser)) -> Any: sort: Optional[str] = "updated_at",
_: User = Depends(get_current_active_superuser),
) -> Any:
""" """
查询当前目录下所有目录和文件 查询当前目录下所有目录和文件
:param fileitem: 文件项 :param fileitem: 文件项
@@ -99,9 +107,11 @@ def list_files(fileitem: schemas.FileItem,
@router.post("/mkdir", summary="创建目录", response_model=schemas.Response) @router.post("/mkdir", summary="创建目录", response_model=schemas.Response)
def mkdir(fileitem: schemas.FileItem, def mkdir(
fileitem: schemas.FileItem,
name: str, name: str,
_: User = Depends(get_current_active_superuser)) -> Any: _: User = Depends(get_current_active_superuser),
) -> Any:
""" """
创建目录 创建目录
:param fileitem: 文件项 :param fileitem: 文件项
@@ -117,8 +127,9 @@ def mkdir(fileitem: schemas.FileItem,
@router.post("/delete", summary="删除文件或目录", response_model=schemas.Response) @router.post("/delete", summary="删除文件或目录", response_model=schemas.Response)
def delete(fileitem: schemas.FileItem, def delete(
_: User = Depends(get_current_active_superuser)) -> Any: fileitem: schemas.FileItem, _: User = Depends(get_current_active_superuser)
) -> Any:
""" """
删除文件或目录 删除文件或目录
:param fileitem: 文件项 :param fileitem: 文件项
@@ -131,8 +142,9 @@ def delete(fileitem: schemas.FileItem,
@router.post("/download", summary="下载文件") @router.post("/download", summary="下载文件")
def download(fileitem: schemas.FileItem, def download(
_: User = Depends(get_current_active_superuser)) -> Any: fileitem: schemas.FileItem, _: User = Depends(get_current_active_superuser)
) -> Any:
""" """
下载文件或目录 下载文件或目录
:param fileitem: 文件项 :param fileitem: 文件项
@@ -146,8 +158,9 @@ def download(fileitem: schemas.FileItem,
@router.post("/image", summary="预览图片") @router.post("/image", summary="预览图片")
def image(fileitem: schemas.FileItem, def image(
_: User = Depends(get_current_active_superuser)) -> Any: fileitem: schemas.FileItem, _: User = Depends(get_current_active_superuser)
) -> Any:
""" """
下载文件或目录 下载文件或目录
:param fileitem: 文件项 :param fileitem: 文件项
@@ -161,10 +174,12 @@ def image(fileitem: schemas.FileItem,
@router.post("/rename", summary="重命名文件或目录", response_model=schemas.Response) @router.post("/rename", summary="重命名文件或目录", response_model=schemas.Response)
def rename(fileitem: schemas.FileItem, def rename(
fileitem: schemas.FileItem,
new_name: str, new_name: str,
recursive: Optional[bool] = False, recursive: Optional[bool] = False,
_: User = Depends(get_current_active_superuser)) -> Any: _: User = Depends(get_current_active_superuser),
) -> Any:
""" """
重命名文件或目录 重命名文件或目录
:param fileitem: 文件项 :param fileitem: 文件项
@@ -189,8 +204,9 @@ def rename(fileitem: schemas.FileItem,
handled = 0 handled = 0
for sub_file in sub_files: for sub_file in sub_files:
handled += 1 handled += 1
progress.update(value=handled / total * 100, progress.update(
text=f"正在处理 {sub_file.name} ...") value=handled / total * 100, text=f"正在处理 {sub_file.name} ..."
)
if sub_file.type == "dir": if sub_file.type == "dir":
continue continue
if not sub_file.extension: if not sub_file.extension:
@@ -204,20 +220,25 @@ def rename(fileitem: schemas.FileItem,
) )
if not context or not context.media_info: if not context or not context.media_info:
progress.end() progress.end()
return schemas.Response(success=False, message=f"{sub_path.name} 未识别到媒体信息") return schemas.Response(
success=False, message=f"{sub_path.name} 未识别到媒体信息"
)
new_path = transferchain.recommend_name( new_path = transferchain.recommend_name(
meta=context.meta_info, meta=context.meta_info, mediainfo=context.media_info
mediainfo=context.media_info
) )
if not new_path: if not new_path:
progress.end() progress.end()
return schemas.Response(success=False, message=f"{sub_path.name} 未识别到新名称") return schemas.Response(
ret: schemas.Response = rename(fileitem=sub_file, success=False, message=f"{sub_path.name} 未识别到新名称"
new_name=Path(new_path).name, )
recursive=False) ret: schemas.Response = rename(
fileitem=sub_file, new_name=Path(new_path).name, recursive=False
)
if not ret.success: if not ret.success:
progress.end() progress.end()
return schemas.Response(success=False, message=f"{sub_path.name} 重命名失败!") return schemas.Response(
success=False, message=f"{sub_path.name} 重命名失败!"
)
progress.end() progress.end()
# 重命名自己 # 重命名自己
result = StorageChain().rename_file(fileitem, new_name) result = StorageChain().rename_file(fileitem, new_name)
@@ -226,7 +247,9 @@ def rename(fileitem: schemas.FileItem,
return schemas.Response(success=False) return schemas.Response(success=False)
@router.get("/usage/{name}", summary="存储空间信息", response_model=schemas.StorageUsage) @router.get(
"/usage/{name}", summary="存储空间信息", response_model=schemas.StorageUsage
)
def usage(name: str, _: User = Depends(get_current_active_superuser)) -> Any: def usage(name: str, _: User = Depends(get_current_active_superuser)) -> Any:
""" """
查询存储空间 查询存储空间
@@ -237,8 +260,14 @@ def usage(name: str, _: User = Depends(get_current_active_superuser)) -> Any:
return schemas.StorageUsage() return schemas.StorageUsage()
@router.get("/transtype/{name}", summary="支持的整理方式获取", response_model=schemas.StorageTransType) @router.get(
async def transtype(name: str, _: User = Depends(get_current_active_superuser_async)) -> Any: "/transtype/{name}",
summary="支持的整理方式获取",
response_model=schemas.StorageTransType,
)
async def transtype(
name: str, _: User = Depends(get_current_active_superuser_async)
) -> Any:
""" """
查询支持的整理方式 查询支持的整理方式
""" """
+183 -110
View File
@@ -25,26 +25,36 @@ from app.schemas.types import MediaType, EventType, SystemConfigKey
router = APIRouter() router = APIRouter()
def start_subscribe_add(title: str, year: str, def start_subscribe_add(
mtype: MediaType, tmdbid: int, season: int, username: str): title: str, year: str, mtype: MediaType, tmdbid: int, season: int, username: str
):
""" """
启动订阅任务 启动订阅任务
""" """
SubscribeChain().add(title=title, year=year, SubscribeChain().add(
mtype=mtype, tmdbid=tmdbid, season=season, username=username) title=title,
year=year,
mtype=mtype,
tmdbid=tmdbid,
season=season,
username=username,
)
@router.get("/", summary="查询所有订阅", response_model=List[schemas.Subscribe]) @router.get("/", summary="查询所有订阅", response_model=List[schemas.Subscribe])
async def read_subscribes( async def read_subscribes(
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
查询所有订阅 查询所有订阅
""" """
return await Subscribe.async_list(db) return await Subscribe.async_list(db)
@router.get("/list", summary="查询所有订阅(API_TOKEN", response_model=List[schemas.Subscribe]) @router.get(
"/list", summary="查询所有订阅(API_TOKEN", response_model=List[schemas.Subscribe]
)
async def list_subscribes(_: Annotated[str, Depends(verify_apitoken)]) -> Any: async def list_subscribes(_: Annotated[str, Depends(verify_apitoken)]) -> Any:
""" """
查询所有订阅 API_TOKEN认证(?token=xxx 查询所有订阅 API_TOKEN认证(?token=xxx
@@ -82,13 +92,10 @@ async def create_subscribe(
subscribe_dict = subscribe_in.model_dump() subscribe_dict = subscribe_in.model_dump()
if subscribe_in.id: if subscribe_in.id:
subscribe_dict.pop("id", None) subscribe_dict.pop("id", None)
sid, message = await SubscribeChain().async_add(mtype=mtype, sid, message = await SubscribeChain().async_add(
title=title, mtype=mtype, title=title, exist_ok=True, **subscribe_dict
exist_ok=True,
**subscribe_dict)
return schemas.Response(
success=bool(sid), message=message, data={"id": sid}
) )
return schemas.Response(success=bool(sid), message=message, data={"id": sid})
@router.put("/", summary="更新订阅", response_model=schemas.Response) @router.put("/", summary="更新订阅", response_model=schemas.Response)
@@ -96,7 +103,7 @@ async def update_subscribe(
*, *,
subscribe_in: schemas.Subscribe, subscribe_in: schemas.Subscribe,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token) _: schemas.TokenPayload = Depends(verify_token),
) -> Any: ) -> Any:
""" """
更新订阅信息 更新订阅信息
@@ -115,9 +122,9 @@ async def update_subscribe(
elif subscribe_in.total_episode: elif subscribe_in.total_episode:
# 总集数增加时,缺失集数也要增加 # 总集数增加时,缺失集数也要增加
if subscribe_in.total_episode > (subscribe.total_episode or 0): if subscribe_in.total_episode > (subscribe.total_episode or 0):
subscribe_dict["lack_episode"] = (subscribe.lack_episode subscribe_dict["lack_episode"] = subscribe.lack_episode + (
+ (subscribe_in.total_episode subscribe_in.total_episode - (subscribe.total_episode or 0)
- (subscribe.total_episode or 0))) )
# 是否手动修改过总集数 # 是否手动修改过总集数
if subscribe_in.total_episode != subscribe.total_episode: if subscribe_in.total_episode != subscribe.total_episode:
subscribe_dict["manual_total_episode"] = 1 subscribe_dict["manual_total_episode"] = 1
@@ -126,11 +133,14 @@ async def update_subscribe(
# 重新获取更新后的订阅数据 # 重新获取更新后的订阅数据
updated_subscribe = await Subscribe.async_get(db, subscribe_in.id) updated_subscribe = await Subscribe.async_get(db, subscribe_in.id)
# 发送订阅调整事件 # 发送订阅调整事件
await eventmanager.async_send_event(EventType.SubscribeModified, { await eventmanager.async_send_event(
EventType.SubscribeModified,
{
"subscribe_id": subscribe_in.id, "subscribe_id": subscribe_in.id,
"old_subscribe_info": old_subscribe_dict, "old_subscribe_info": old_subscribe_dict,
"subscribe_info": updated_subscribe.to_dict() if updated_subscribe else {}, "subscribe_info": updated_subscribe.to_dict() if updated_subscribe else {},
}) },
)
return schemas.Response(success=True) return schemas.Response(success=True)
@@ -139,7 +149,8 @@ async def update_subscribe_status(
subid: int, subid: int,
state: str, state: str,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
更新订阅状态 更新订阅状态
""" """
@@ -150,17 +161,18 @@ async def update_subscribe_status(
if state not in valid_states: if state not in valid_states:
return schemas.Response(success=False, message="无效的订阅状态") return schemas.Response(success=False, message="无效的订阅状态")
old_subscribe_dict = subscribe.to_dict() old_subscribe_dict = subscribe.to_dict()
await subscribe.async_update(db, { await subscribe.async_update(db, {"state": state})
"state": state
})
# 重新获取更新后的订阅数据 # 重新获取更新后的订阅数据
updated_subscribe = await Subscribe.async_get(db, subid) updated_subscribe = await Subscribe.async_get(db, subid)
# 发送订阅调整事件 # 发送订阅调整事件
await eventmanager.async_send_event(EventType.SubscribeModified, { await eventmanager.async_send_event(
EventType.SubscribeModified,
{
"subscribe_id": subid, "subscribe_id": subid,
"old_subscribe_info": old_subscribe_dict, "old_subscribe_info": old_subscribe_dict,
"subscribe_info": updated_subscribe.to_dict() if updated_subscribe else {}, "subscribe_info": updated_subscribe.to_dict() if updated_subscribe else {},
}) },
)
return schemas.Response(success=True) return schemas.Response(success=True)
@@ -170,7 +182,8 @@ async def subscribe_mediaid(
season: Optional[int] = None, season: Optional[int] = None,
title: Optional[str] = None, title: Optional[str] = None,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
根据 TMDBID/豆瓣ID/BangumiId 查询订阅 tmdb:/douban: 根据 TMDBID/豆瓣ID/BangumiId 查询订阅 tmdb:/douban:
""" """
@@ -203,14 +216,15 @@ async def subscribe_mediaid(
meta = MetaInfo(title) meta = MetaInfo(title)
if season is not None: if season is not None:
meta.begin_season = season meta.begin_season = season
result = await Subscribe.async_get_by_title(db, title=meta.name, season=meta.begin_season) result = await Subscribe.async_get_by_title(
db, title=meta.name, season=meta.begin_season
)
return result if result else Subscribe() return result if result else Subscribe()
@router.get("/refresh", summary="刷新订阅", response_model=schemas.Response) @router.get("/refresh", summary="刷新订阅", response_model=schemas.Response)
def refresh_subscribes( def refresh_subscribes(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
_: schemas.TokenPayload = Depends(verify_token)) -> Any:
""" """
刷新所有订阅 刷新所有订阅
""" """
@@ -222,7 +236,8 @@ def refresh_subscribes(
async def reset_subscribes( async def reset_subscribes(
subid: int, subid: int,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
重置订阅 重置订阅
""" """
@@ -231,28 +246,35 @@ async def reset_subscribes(
# 在更新之前获取旧数据 # 在更新之前获取旧数据
old_subscribe_dict = subscribe.to_dict() old_subscribe_dict = subscribe.to_dict()
# 更新订阅 # 更新订阅
await subscribe.async_update(db, { await subscribe.async_update(
db,
{
"note": [], "note": [],
"lack_episode": subscribe.total_episode, "lack_episode": subscribe.total_episode,
"current_priority": None, "current_priority": None,
"episode_priority": {}, "episode_priority": {},
"state": "R" "state": "R",
}) },
)
# 重新获取更新后的订阅数据 # 重新获取更新后的订阅数据
updated_subscribe = await Subscribe.async_get(db, subid) updated_subscribe = await Subscribe.async_get(db, subid)
# 发送订阅调整事件 # 发送订阅调整事件
await eventmanager.async_send_event(EventType.SubscribeModified, { await eventmanager.async_send_event(
EventType.SubscribeModified,
{
"subscribe_id": subid, "subscribe_id": subid,
"old_subscribe_info": old_subscribe_dict, "old_subscribe_info": old_subscribe_dict,
"subscribe_info": updated_subscribe.to_dict() if updated_subscribe else {}, "subscribe_info": updated_subscribe.to_dict()
}) if updated_subscribe
else {},
},
)
return schemas.Response(success=True) return schemas.Response(success=True)
return schemas.Response(success=False, message="订阅不存在") return schemas.Response(success=False, message="订阅不存在")
@router.get("/check", summary="刷新订阅 TMDB 信息", response_model=schemas.Response) @router.get("/check", summary="刷新订阅 TMDB 信息", response_model=schemas.Response)
def check_subscribes( def check_subscribes(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
_: schemas.TokenPayload = Depends(verify_token)) -> Any:
""" """
刷新订阅 TMDB 信息 刷新订阅 TMDB 信息
""" """
@@ -262,39 +284,34 @@ def check_subscribes(
@router.get("/search", summary="搜索所有订阅", response_model=schemas.Response) @router.get("/search", summary="搜索所有订阅", response_model=schemas.Response)
async def search_subscribes( async def search_subscribes(
background_tasks: BackgroundTasks, background_tasks: BackgroundTasks, _: schemas.TokenPayload = Depends(verify_token)
_: schemas.TokenPayload = Depends(verify_token)) -> Any: ) -> Any:
""" """
搜索所有订阅 搜索所有订阅
""" """
background_tasks.add_task( background_tasks.add_task(
Scheduler().start, Scheduler().start,
job_id="subscribe_search", job_id="subscribe_search",
**{ **{"sid": None, "state": "R", "manual": True},
"sid": None,
"state": 'R',
"manual": True
}
) )
return schemas.Response(success=True) return schemas.Response(success=True)
@router.get("/search/{subscribe_id}", summary="搜索订阅", response_model=schemas.Response) @router.get(
"/search/{subscribe_id}", summary="搜索订阅", response_model=schemas.Response
)
async def search_subscribe( async def search_subscribe(
subscribe_id: int, subscribe_id: int,
background_tasks: BackgroundTasks, background_tasks: BackgroundTasks,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
根据订阅编号搜索订阅 根据订阅编号搜索订阅
""" """
background_tasks.add_task( background_tasks.add_task(
Scheduler().start, Scheduler().start,
job_id="subscribe_search", job_id="subscribe_search",
**{ **{"sid": subscribe_id, "state": None, "manual": True},
"sid": subscribe_id,
"state": None,
"manual": True
}
) )
return schemas.Response(success=True) return schemas.Response(success=True)
@@ -304,7 +321,7 @@ async def delete_subscribe_by_mediaid(
mediaid: str, mediaid: str,
season: Optional[int] = None, season: Optional[int] = None,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token) _: schemas.TokenPayload = Depends(verify_token),
) -> Any: ) -> Any:
""" """
根据TMDBID或豆瓣ID删除订阅 tmdb:/douban: 根据TMDBID或豆瓣ID删除订阅 tmdb:/douban:
@@ -333,16 +350,21 @@ async def delete_subscribe_by_mediaid(
subscribe_id = subscribe.id subscribe_id = subscribe.id
await Subscribe.async_delete(db, subscribe_id) await Subscribe.async_delete(db, subscribe_id)
# 发送事件 # 发送事件
await eventmanager.async_send_event(EventType.SubscribeDeleted, { await eventmanager.async_send_event(
"subscribe_id": subscribe_id, EventType.SubscribeDeleted,
"subscribe_info": subscribe_info {"subscribe_id": subscribe_id, "subscribe_info": subscribe_info},
}) )
return schemas.Response(success=True) return schemas.Response(success=True)
@router.post("/seerr", summary="OverSeerr/JellySeerr通知订阅", response_model=schemas.Response) @router.post(
async def seerr_subscribe(request: Request, background_tasks: BackgroundTasks, "/seerr", summary="OverSeerr/JellySeerr通知订阅", response_model=schemas.Response
authorization: Annotated[str | None, Header()] = None) -> Any: )
async def seerr_subscribe(
request: Request,
background_tasks: BackgroundTasks,
authorization: Annotated[str | None, Header()] = None,
) -> Any:
""" """
Jellyseerr/Overseerr网络勾子通知订阅 Jellyseerr/Overseerr网络勾子通知订阅
""" """
@@ -361,49 +383,66 @@ async def seerr_subscribe(request: Request, background_tasks: BackgroundTasks,
if notification_type not in ["MEDIA_APPROVED", "MEDIA_AUTO_APPROVED"]: if notification_type not in ["MEDIA_APPROVED", "MEDIA_AUTO_APPROVED"]:
return schemas.Response(success=False, message="不支持的通知类型") return schemas.Response(success=False, message="不支持的通知类型")
subject = req_json.get("subject") subject = req_json.get("subject")
media_type = MediaType.MOVIE if req_json.get("media", {}).get("media_type") == "movie" else MediaType.TV media_type = (
MediaType.MOVIE
if req_json.get("media", {}).get("media_type") == "movie"
else MediaType.TV
)
tmdbId = req_json.get("media", {}).get("tmdbId") tmdbId = req_json.get("media", {}).get("tmdbId")
if not media_type or not tmdbId or not subject: if not media_type or not tmdbId or not subject:
return schemas.Response(success=False, message="请求参数不正确") return schemas.Response(success=False, message="请求参数不正确")
user_name = req_json.get("request", {}).get("requestedBy_username") user_name = req_json.get("request", {}).get("requestedBy_username")
# 添加订阅 # 添加订阅
if media_type == MediaType.MOVIE: if media_type == MediaType.MOVIE:
background_tasks.add_task(start_subscribe_add, background_tasks.add_task(
start_subscribe_add,
mtype=media_type, mtype=media_type,
tmdbid=tmdbId, tmdbid=tmdbId,
title=subject, title=subject,
year="", year="",
season=0, season=0,
username=user_name) username=user_name,
)
else: else:
seasons = [] seasons = []
for extra in req_json.get("extra", []): for extra in req_json.get("extra", []):
if extra.get("name") == "Requested Seasons": if extra.get("name") == "Requested Seasons":
seasons = [int(str(sea).strip()) for sea in extra.get("value").split(", ") if str(sea).isdigit()] seasons = [
int(str(sea).strip())
for sea in extra.get("value").split(", ")
if str(sea).isdigit()
]
break break
for season in seasons: for season in seasons:
background_tasks.add_task(start_subscribe_add, background_tasks.add_task(
start_subscribe_add,
mtype=media_type, mtype=media_type,
tmdbid=tmdbId, tmdbid=tmdbId,
title=subject, title=subject,
year="", year="",
season=season, season=season,
username=user_name) username=user_name,
)
return schemas.Response(success=True) return schemas.Response(success=True)
@router.get("/history/{mtype}", summary="查询订阅历史", response_model=List[schemas.Subscribe]) @router.get(
"/history/{mtype}", summary="查询订阅历史", response_model=List[schemas.Subscribe]
)
async def subscribe_history( async def subscribe_history(
mtype: str, mtype: str,
page: Optional[int] = 1, page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
查询电影/电视剧订阅历史 查询电影/电视剧订阅历史
""" """
histories = await SubscribeHistory.async_list_by_type(db, mtype=mtype, page=page, count=count) histories = await SubscribeHistory.async_list_by_type(
db, mtype=mtype, page=page, count=count
)
result = [] result = []
for history in histories: for history in histories:
history_item = schemas.Subscribe.model_validate(history, from_attributes=True) history_item = schemas.Subscribe.model_validate(history, from_attributes=True)
@@ -414,11 +453,13 @@ async def subscribe_history(
return result return result
@router.delete("/history/{history_id}", summary="删除订阅历史", response_model=schemas.Response) @router.delete(
async def delete_subscribe( "/history/{history_id}", summary="删除订阅历史", response_model=schemas.Response
)
async def delete_subscribe_history(
history_id: int, history_id: int,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token) _: schemas.TokenPayload = Depends(verify_token),
) -> Any: ) -> Any:
""" """
删除订阅历史 删除订阅历史
@@ -427,7 +468,11 @@ async def delete_subscribe(
return schemas.Response(success=True) return schemas.Response(success=True)
@router.get("/popular", summary="热门订阅(基于用户共享数据)", response_model=List[schemas.MediaInfo]) @router.get(
"/popular",
summary="热门订阅(基于用户共享数据)",
response_model=List[schemas.MediaInfo],
)
async def popular_subscribes( async def popular_subscribes(
stype: str, stype: str,
page: Optional[int] = 1, page: Optional[int] = 1,
@@ -437,7 +482,8 @@ async def popular_subscribes(
min_rating: Optional[float] = None, min_rating: Optional[float] = None,
max_rating: Optional[float] = None, max_rating: Optional[float] = None,
sort_type: Optional[str] = None, sort_type: Optional[str] = None,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
查询热门订阅 查询热门订阅
""" """
@@ -448,7 +494,7 @@ async def popular_subscribes(
genre_id=genre_id, genre_id=genre_id,
min_rating=min_rating, min_rating=min_rating,
max_rating=max_rating, max_rating=max_rating,
sort_type=sort_type sort_type=sort_type,
) )
if subscribes: if subscribes:
ret_medias = [] ret_medias = []
@@ -484,22 +530,30 @@ async def popular_subscribes(
return [] return []
@router.get("/user/{username}", summary="用户订阅", response_model=List[schemas.Subscribe]) @router.get(
"/user/{username}", summary="用户订阅", response_model=List[schemas.Subscribe]
)
async def user_subscribes( async def user_subscribes(
username: str, username: str,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
查询用户订阅 查询用户订阅
""" """
return await Subscribe.async_list_by_username(db, username) return await Subscribe.async_list_by_username(db, username)
@router.get("/files/{subscribe_id}", summary="订阅相关文件信息", response_model=schemas.SubscrbieInfo) @router.get(
"/files/{subscribe_id}",
summary="订阅相关文件信息",
response_model=schemas.SubscrbieInfo,
)
def subscribe_files( def subscribe_files(
subscribe_id: int, subscribe_id: int,
db: Session = Depends(get_db), db: Session = Depends(get_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
订阅相关文件信息 订阅相关文件信息
""" """
@@ -511,22 +565,24 @@ def subscribe_files(
@router.post("/share", summary="分享订阅", response_model=schemas.Response) @router.post("/share", summary="分享订阅", response_model=schemas.Response)
async def subscribe_share( async def subscribe_share(
sub: schemas.SubscribeShare, sub: schemas.SubscribeShare, _: schemas.TokenPayload = Depends(verify_token)
_: schemas.TokenPayload = Depends(verify_token)) -> Any: ) -> Any:
""" """
分享订阅 分享订阅
""" """
state, errmsg = await SubscribeHelper().async_sub_share(subscribe_id=sub.subscribe_id, state, errmsg = await SubscribeHelper().async_sub_share(
subscribe_id=sub.subscribe_id,
share_title=sub.share_title, share_title=sub.share_title,
share_comment=sub.share_comment, share_comment=sub.share_comment,
share_user=sub.share_user) share_user=sub.share_user,
)
return schemas.Response(success=state, message=errmsg) return schemas.Response(success=state, message=errmsg)
@router.delete("/share/{share_id}", summary="删除分享", response_model=schemas.Response) @router.delete("/share/{share_id}", summary="删除分享", response_model=schemas.Response)
async def subscribe_share_delete( async def subscribe_share_delete(
share_id: int, share_id: int, _: schemas.TokenPayload = Depends(verify_token)
_: schemas.TokenPayload = Depends(verify_token)) -> Any: ) -> Any:
""" """
删除分享 删除分享
""" """
@@ -537,7 +593,8 @@ async def subscribe_share_delete(
@router.post("/fork", summary="复用订阅", response_model=schemas.Response) @router.post("/fork", summary="复用订阅", response_model=schemas.Response)
async def subscribe_fork( async def subscribe_fork(
sub: schemas.SubscribeShare, sub: schemas.SubscribeShare,
current_user: User = Depends(get_current_active_user_async)) -> Any: current_user: User = Depends(get_current_active_user_async),
) -> Any:
""" """
复用订阅 复用订阅
""" """
@@ -546,8 +603,9 @@ async def subscribe_fork(
for key in list(sub_dict.keys()): for key in list(sub_dict.keys()):
if not hasattr(schemas.Subscribe(), key): if not hasattr(schemas.Subscribe(), key):
sub_dict.pop(key) sub_dict.pop(key)
result = await create_subscribe(subscribe_in=schemas.Subscribe(**sub_dict), result = await create_subscribe(
current_user=current_user) subscribe_in=schemas.Subscribe(**sub_dict), current_user=current_user
)
if result.success: if result.success:
await SubscribeHelper().async_sub_fork(share_id=sub.id) await SubscribeHelper().async_sub_fork(share_id=sub.id)
return result return result
@@ -563,34 +621,42 @@ async def followed_subscribers(_: schemas.TokenPayload = Depends(verify_token))
@router.post("/follow", summary="Follow订阅分享人", response_model=schemas.Response) @router.post("/follow", summary="Follow订阅分享人", response_model=schemas.Response)
async def follow_subscriber( async def follow_subscriber(
share_uid: Optional[str] = None, share_uid: Optional[str] = None, _: schemas.TokenPayload = Depends(verify_token)
_: schemas.TokenPayload = Depends(verify_token)) -> Any: ) -> Any:
""" """
Follow订阅分享人 Follow订阅分享人
""" """
subscribers = SystemConfigOper().get(SystemConfigKey.FollowSubscribers) or [] subscribers = SystemConfigOper().get(SystemConfigKey.FollowSubscribers) or []
if share_uid and share_uid not in subscribers: if share_uid and share_uid not in subscribers:
subscribers.append(share_uid) subscribers.append(share_uid)
await SystemConfigOper().async_set(SystemConfigKey.FollowSubscribers, subscribers) await SystemConfigOper().async_set(
SystemConfigKey.FollowSubscribers, subscribers
)
return schemas.Response(success=True) return schemas.Response(success=True)
@router.delete("/follow", summary="取消Follow订阅分享人", response_model=schemas.Response) @router.delete(
"/follow", summary="取消Follow订阅分享人", response_model=schemas.Response
)
async def unfollow_subscriber( async def unfollow_subscriber(
share_uid: Optional[str] = None, share_uid: Optional[str] = None, _: schemas.TokenPayload = Depends(verify_token)
_: schemas.TokenPayload = Depends(verify_token)) -> Any: ) -> Any:
""" """
取消Follow订阅分享人 取消Follow订阅分享人
""" """
subscribers = SystemConfigOper().get(SystemConfigKey.FollowSubscribers) or [] subscribers = SystemConfigOper().get(SystemConfigKey.FollowSubscribers) or []
if share_uid and share_uid in subscribers: if share_uid and share_uid in subscribers:
subscribers.remove(share_uid) subscribers.remove(share_uid)
await SystemConfigOper().async_set(SystemConfigKey.FollowSubscribers, subscribers) await SystemConfigOper().async_set(
SystemConfigKey.FollowSubscribers, subscribers
)
return schemas.Response(success=True) return schemas.Response(success=True)
@router.get("/shares", summary="查询分享的订阅", response_model=List[schemas.SubscribeShare]) @router.get(
async def popular_subscribes( "/shares", summary="查询分享的订阅", response_model=List[schemas.SubscribeShare]
)
async def subscribe_shares(
name: Optional[str] = None, name: Optional[str] = None,
page: Optional[int] = 1, page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
@@ -598,7 +664,8 @@ async def popular_subscribes(
min_rating: Optional[float] = None, min_rating: Optional[float] = None,
max_rating: Optional[float] = None, max_rating: Optional[float] = None,
sort_type: Optional[str] = None, sort_type: Optional[str] = None,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
查询分享的订阅 查询分享的订阅
""" """
@@ -609,12 +676,18 @@ async def popular_subscribes(
genre_id=genre_id, genre_id=genre_id,
min_rating=min_rating, min_rating=min_rating,
max_rating=max_rating, max_rating=max_rating,
sort_type=sort_type sort_type=sort_type,
) )
@router.get("/share/statistics", summary="查询订阅分享统计", response_model=List[schemas.SubscribeShareStatistics]) @router.get(
async def subscribe_share_statistics(_: schemas.TokenPayload = Depends(verify_token)) -> Any: "/share/statistics",
summary="查询订阅分享统计",
response_model=List[schemas.SubscribeShareStatistics],
)
async def subscribe_share_statistics(
_: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
查询订阅分享统计 查询订阅分享统计
返回每个分享人分享的媒体数量以及总的复用人次 返回每个分享人分享的媒体数量以及总的复用人次
@@ -626,7 +699,8 @@ async def subscribe_share_statistics(_: schemas.TokenPayload = Depends(verify_to
async def read_subscribe( async def read_subscribe(
subscribe_id: int, subscribe_id: int,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
根据订阅编号查询订阅信息 根据订阅编号查询订阅信息
""" """
@@ -639,7 +713,7 @@ async def read_subscribe(
async def delete_subscribe( async def delete_subscribe(
subscribe_id: int, subscribe_id: int,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token) _: schemas.TokenPayload = Depends(verify_token),
) -> Any: ) -> Any:
""" """
删除订阅信息 删除订阅信息
@@ -650,13 +724,12 @@ async def delete_subscribe(
subscribe_info = subscribe.to_dict() subscribe_info = subscribe.to_dict()
await Subscribe.async_delete(db, subscribe_id) await Subscribe.async_delete(db, subscribe_id)
# 发送事件 # 发送事件
await eventmanager.async_send_event(EventType.SubscribeDeleted, { await eventmanager.async_send_event(
"subscribe_id": subscribe_id, EventType.SubscribeDeleted,
"subscribe_info": subscribe_info {"subscribe_id": subscribe_id, "subscribe_info": subscribe_info},
}) )
# 统计订阅 # 统计订阅
SubscribeHelper().sub_done_async({ SubscribeHelper().sub_done_async(
"tmdbid": subscribe.tmdbid, {"tmdbid": subscribe.tmdbid, "doubanid": subscribe.doubanid}
"doubanid": subscribe.doubanid )
})
return schemas.Response(success=True) return schemas.Response(success=True)
+27 -8
View File
@@ -65,7 +65,9 @@ def _match_nettest_prefix(url: str, prefix: str) -> bool:
if (parsed_url.hostname or "").lower() != (parsed_prefix.hostname or "").lower(): if (parsed_url.hostname or "").lower() != (parsed_prefix.hostname or "").lower():
return False return False
url_port = parsed_url.port or (443 if parsed_url.scheme.lower() == "https" else 80) url_port = parsed_url.port or (443 if parsed_url.scheme.lower() == "https" else 80)
prefix_port = parsed_prefix.port or (443 if parsed_prefix.scheme.lower() == "https" else 80) prefix_port = parsed_prefix.port or (
443 if parsed_prefix.scheme.lower() == "https" else 80
)
if url_port != prefix_port: if url_port != prefix_port:
return False return False
return parsed_url.path.startswith(parsed_prefix.path or "/") return parsed_url.path.startswith(parsed_prefix.path or "/")
@@ -193,14 +195,18 @@ def _build_nettest_rules() -> list[dict[str, Any]]:
"id": "github_proxy_web", "id": "github_proxy_web",
"name": "github.com", "name": "github.com",
"icon": "github", "icon": "github",
"url": f"{github_proxy}{github_readme_url}" if github_proxy else github_readme_url, "url": f"{github_proxy}{github_readme_url}"
if github_proxy
else github_readme_url,
"proxy": True, "proxy": True,
"allowed_redirect_prefixes": [ "allowed_redirect_prefixes": [
"https://github.com/", "https://github.com/",
*((f"{github_proxy}https://github.com/",) if github_proxy else ()), *((f"{github_proxy}https://github.com/",) if github_proxy else ()),
], ],
"expected_text": "MoviePilot", "expected_text": "MoviePilot",
"invalid_message": "Github加速代理已失效,请检查配置" if github_proxy else "无效响应", "invalid_message": "Github加速代理已失效,请检查配置"
if github_proxy
else "无效响应",
"proxy_name": "Github加速代理" if github_proxy else "", "proxy_name": "Github加速代理" if github_proxy else "",
"headers": settings.GITHUB_HEADERS, "headers": settings.GITHUB_HEADERS,
}, },
@@ -229,14 +235,22 @@ def _build_nettest_rules() -> list[dict[str, Any]]:
"id": "github_proxy_raw", "id": "github_proxy_raw",
"name": "raw.githubusercontent.com", "name": "raw.githubusercontent.com",
"icon": "github", "icon": "github",
"url": f"{github_proxy}{raw_readme_url}" if github_proxy else raw_readme_url, "url": f"{github_proxy}{raw_readme_url}"
if github_proxy
else raw_readme_url,
"proxy": True, "proxy": True,
"allowed_redirect_prefixes": [ "allowed_redirect_prefixes": [
"https://raw.githubusercontent.com/", "https://raw.githubusercontent.com/",
*((f"{github_proxy}https://raw.githubusercontent.com/",) if github_proxy else ()), *(
(f"{github_proxy}https://raw.githubusercontent.com/",)
if github_proxy
else ()
),
], ],
"expected_text": "MoviePilot", "expected_text": "MoviePilot",
"invalid_message": "Github加速代理已失效,请检查配置" if github_proxy else "无效响应", "invalid_message": "Github加速代理已失效,请检查配置"
if github_proxy
else "无效响应",
"proxy_name": "Github加速代理" if github_proxy else "", "proxy_name": "Github加速代理" if github_proxy else "",
"headers": settings.GITHUB_HEADERS, "headers": settings.GITHUB_HEADERS,
}, },
@@ -257,6 +271,7 @@ def _build_nettest_rules() -> list[dict[str, Any]]:
) )
return rules return rules
def _validate_nettest_url(url: str) -> Optional[str]: def _validate_nettest_url(url: str) -> Optional[str]:
""" """
对实际请求地址做基础安全校验。 对实际请求地址做基础安全校验。
@@ -276,7 +291,9 @@ def _validate_nettest_url(url: str) -> Optional[str]:
return None return None
def _get_nettest_rule(url: Optional[str] = None, target_id: Optional[str] = None) -> Optional[dict[str, Any]]: def _get_nettest_rule(
url: Optional[str] = None, target_id: Optional[str] = None
) -> Optional[dict[str, Any]]:
""" """
根据 target_id 或历史兼容参数匹配网络测试规则。 根据 target_id 或历史兼容参数匹配网络测试规则。
@@ -804,7 +821,9 @@ def ruletest(
) )
@router.get("/nettest/targets", summary="获取网络测试目标", response_model=schemas.Response) @router.get(
"/nettest/targets", summary="获取网络测试目标", response_model=schemas.Response
)
async def nettest_targets(_: schemas.TokenPayload = Depends(verify_token)): async def nettest_targets(_: schemas.TokenPayload = Depends(verify_token)):
""" """
获取网络测试目标。 获取网络测试目标。
+69 -26
View File
@@ -10,8 +10,12 @@ from app.schemas.types import MediaType
router = APIRouter() router = APIRouter()
@router.get("/seasons/{tmdbid}", summary="TMDB所有季", response_model=List[schemas.TmdbSeason]) @router.get(
async def tmdb_seasons(tmdbid: int, _: schemas.TokenPayload = Depends(verify_token)) -> Any: "/seasons/{tmdbid}", summary="TMDB所有季", response_model=List[schemas.TmdbSeason]
)
async def tmdb_seasons(
tmdbid: int, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
根据TMDBID查询themoviedb所有季信息 根据TMDBID查询themoviedb所有季信息
""" """
@@ -21,10 +25,14 @@ async def tmdb_seasons(tmdbid: int, _: schemas.TokenPayload = Depends(verify_tok
return [] return []
@router.get("/similar/{tmdbid}/{type_name}", summary="类似电影/电视剧", response_model=List[schemas.MediaInfo]) @router.get(
async def tmdb_similar(tmdbid: int, "/similar/{tmdbid}/{type_name}",
type_name: str, summary="类似电影/电视剧",
_: schemas.TokenPayload = Depends(verify_token)) -> Any: response_model=List[schemas.MediaInfo],
)
async def tmdb_similar(
tmdbid: int, type_name: str, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
根据TMDBID查询类似电影/电视剧,type_name: 电影/电视剧 根据TMDBID查询类似电影/电视剧,type_name: 电影/电视剧
""" """
@@ -40,10 +48,14 @@ async def tmdb_similar(tmdbid: int,
return [] return []
@router.get("/recommend/{tmdbid}/{type_name}", summary="推荐电影/电视剧", response_model=List[schemas.MediaInfo]) @router.get(
async def tmdb_recommend(tmdbid: int, "/recommend/{tmdbid}/{type_name}",
type_name: str, summary="推荐电影/电视剧",
_: schemas.TokenPayload = Depends(verify_token)) -> Any: response_model=List[schemas.MediaInfo],
)
async def tmdb_recommend(
tmdbid: int, type_name: str, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
根据TMDBID查询推荐电影/电视剧,type_name: 电影/电视剧 根据TMDBID查询推荐电影/电视剧,type_name: 电影/电视剧
""" """
@@ -59,11 +71,17 @@ async def tmdb_recommend(tmdbid: int,
return [] return []
@router.get("/collection/{collection_id}", summary="系列合集详情", response_model=List[schemas.MediaInfo]) @router.get(
async def tmdb_collection(collection_id: int, "/collection/{collection_id}",
summary="系列合集详情",
response_model=List[schemas.MediaInfo],
)
async def tmdb_collection(
collection_id: int,
page: Optional[int] = 1, page: Optional[int] = 1,
count: Optional[int] = 20, count: Optional[int] = 20,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
根据合集ID查询合集详情 根据合集ID查询合集详情
""" """
@@ -73,11 +91,17 @@ async def tmdb_collection(collection_id: int,
return [] return []
@router.get("/credits/{tmdbid}/{type_name}", summary="演员阵容", response_model=List[schemas.MediaPerson]) @router.get(
async def tmdb_credits(tmdbid: int, "/credits/{tmdbid}/{type_name}",
summary="演员阵容",
response_model=List[schemas.MediaPerson],
)
async def tmdb_credits(
tmdbid: int,
type_name: str, type_name: str,
page: Optional[int] = 1, page: Optional[int] = 1,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
根据TMDBID查询演员阵容,type_name: 电影/电视剧 根据TMDBID查询演员阵容,type_name: 电影/电视剧
""" """
@@ -91,19 +115,28 @@ async def tmdb_credits(tmdbid: int,
return persons or [] return persons or []
@router.get("/person/{person_id}", summary="人物详情", response_model=schemas.MediaPerson) @router.get(
async def tmdb_person(person_id: int, "/person/{person_id}", summary="人物详情", response_model=schemas.MediaPerson
_: schemas.TokenPayload = Depends(verify_token)) -> Any: )
async def tmdb_person(
person_id: int, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
根据人物ID查询人物详情 根据人物ID查询人物详情
""" """
return await TmdbChain().async_person_detail(person_id=person_id) return await TmdbChain().async_person_detail(person_id=person_id)
@router.get("/person/credits/{person_id}", summary="人物参演作品", response_model=List[schemas.MediaInfo]) @router.get(
async def tmdb_person_credits(person_id: int, "/person/credits/{person_id}",
summary="人物参演作品",
response_model=List[schemas.MediaInfo],
)
async def tmdb_person_credits(
person_id: int,
page: Optional[int] = 1, page: Optional[int] = 1,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
根据人物ID查询人物参演作品 根据人物ID查询人物参演作品
""" """
@@ -113,10 +146,20 @@ async def tmdb_person_credits(person_id: int,
return [] return []
@router.get("/{tmdbid}/{season}", summary="TMDB季所有集", response_model=List[schemas.TmdbEpisode]) @router.get(
async def tmdb_season_episodes(tmdbid: int, season: int, episode_group: Optional[str] = None, "/{tmdbid}/{season}",
_: schemas.TokenPayload = Depends(verify_token)) -> Any: summary="TMDB季所有集",
response_model=List[schemas.TmdbEpisode],
)
async def tmdb_season_episodes(
tmdbid: int,
season: int,
episode_group: Optional[str] = None,
_: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
根据TMDBID查询某季的所有信信息 根据TMDBID查询某季的所有信信息
""" """
return await TmdbChain().async_tmdb_episodes(tmdbid=tmdbid, season=season, episode_group=episode_group) return await TmdbChain().async_tmdb_episodes(
tmdbid=tmdbid, season=season, episode_group=episode_group
)
+84 -29
View File
@@ -9,7 +9,10 @@ from app.core.config import settings
from app.core.context import MediaInfo from app.core.context import MediaInfo
from app.core.metainfo import MetaInfo from app.core.metainfo import MetaInfo
from app.db.models import User from app.db.models import User
from app.db.user_oper import get_current_active_superuser, get_current_active_superuser_async from app.db.user_oper import (
get_current_active_superuser,
get_current_active_superuser_async,
)
from app.utils.crypto import HashUtils from app.utils.crypto import HashUtils
router = APIRouter() router = APIRouter()
@@ -35,8 +38,11 @@ async def torrents_cache(_: User = Depends(get_current_active_superuser_async)):
torrent_data = [] torrent_data = []
for domain, contexts in cache_info.items(): for domain, contexts in cache_info.items():
for context in contexts: for context in contexts:
torrent_hash = HashUtils.md5(f"{context.torrent_info.title}{context.torrent_info.description}") torrent_hash = HashUtils.md5(
torrent_data.append({ f"{context.torrent_info.title}{context.torrent_info.description}"
)
torrent_data.append(
{
"hash": torrent_hash, "hash": torrent_hash,
"domain": domain, "domain": domain,
"title": context.torrent_info.title, "title": context.torrent_info.title,
@@ -44,26 +50,44 @@ async def torrents_cache(_: User = Depends(get_current_active_superuser_async)):
"size": context.torrent_info.size, "size": context.torrent_info.size,
"pubdate": context.torrent_info.pubdate, "pubdate": context.torrent_info.pubdate,
"site_name": context.torrent_info.site_name, "site_name": context.torrent_info.site_name,
"media_name": context.media_info.title if context.media_info else "", "media_name": context.media_info.title
if context.media_info
else "",
"media_year": context.media_info.year if context.media_info else "", "media_year": context.media_info.year if context.media_info else "",
"media_type": context.media_info.type if context.media_info else "", "media_type": context.media_info.type if context.media_info else "",
"season_episode": context.meta_info.season_episode if context.meta_info else "", "season_episode": context.meta_info.season_episode
"resource_term": context.meta_info.resource_term if context.meta_info else "", if context.meta_info
else "",
"resource_term": context.meta_info.resource_term
if context.meta_info
else "",
"enclosure": context.torrent_info.enclosure, "enclosure": context.torrent_info.enclosure,
"page_url": context.torrent_info.page_url, "page_url": context.torrent_info.page_url,
"poster_path": context.media_info.get_poster_image() if context.media_info else "", "poster_path": context.media_info.get_poster_image()
"backdrop_path": context.media_info.get_backdrop_image() if context.media_info else "" if context.media_info
}) else "",
"backdrop_path": context.media_info.get_backdrop_image()
if context.media_info
else "",
}
)
return schemas.Response(success=True, data={ return schemas.Response(
"count": torrent_count, success=True,
"sites": len(cache_info), data={"count": torrent_count, "sites": len(cache_info), "data": torrent_data},
"data": torrent_data )
})
@router.delete("/cache/{domain}/{torrent_hash}", summary="删除指定种子缓存", response_model=schemas.Response) @router.delete(
async def delete_cache(domain: str, torrent_hash: str, _: User = Depends(get_current_active_superuser_async)): "/cache/{domain}/{torrent_hash}",
summary="删除指定种子缓存",
response_model=schemas.Response,
)
async def delete_cache(
domain: str,
torrent_hash: str,
_: User = Depends(get_current_active_superuser_async),
):
""" """
删除指定的种子缓存 删除指定的种子缓存
:param domain: 站点域名 :param domain: 站点域名
@@ -83,8 +107,12 @@ async def delete_cache(domain: str, torrent_hash: str, _: User = Depends(get_cur
# 查找并删除指定种子 # 查找并删除指定种子
original_count = len(cache_data[domain]) original_count = len(cache_data[domain])
cache_data[domain] = [ cache_data[domain] = [
context for context in cache_data[domain] context
if HashUtils.md5(f"{context.torrent_info.title}{context.torrent_info.description}") != torrent_hash for context in cache_data[domain]
if HashUtils.md5(
f"{context.torrent_info.title}{context.torrent_info.description}"
)
!= torrent_hash
] ]
if len(cache_data[domain]) == original_count: if len(cache_data[domain]) == original_count:
@@ -128,15 +156,26 @@ def refresh_cache(_: User = Depends(get_current_active_superuser)):
total_count = sum(len(torrents) for torrents in result.values()) total_count = sum(len(torrents) for torrents in result.values())
sites_count = len(result) sites_count = len(result)
return schemas.Response(success=True, message=f"缓存刷新完成,共刷新 {sites_count} 个站点,{total_count} 个种子") return schemas.Response(
success=True,
message=f"缓存刷新完成,共刷新 {sites_count} 个站点,{total_count} 个种子",
)
except Exception as e: except Exception as e:
return schemas.Response(success=False, message=f"刷新失败:{str(e)}") return schemas.Response(success=False, message=f"刷新失败:{str(e)}")
@router.post("/cache/reidentify/{domain}/{torrent_hash}", summary="重新识别种子", response_model=schemas.Response) @router.post(
async def reidentify_cache(domain: str, torrent_hash: str, "/cache/reidentify/{domain}/{torrent_hash}",
tmdbid: Optional[int] = None, doubanid: Optional[str] = None, summary="重新识别种子",
_: User = Depends(get_current_active_superuser_async)): response_model=schemas.Response,
)
async def reidentify_cache(
domain: str,
torrent_hash: str,
tmdbid: Optional[int] = None,
doubanid: Optional[str] = None,
_: User = Depends(get_current_active_superuser_async),
):
""" """
重新识别指定的种子 重新识别指定的种子
:param domain: 站点域名 :param domain: 站点域名
@@ -159,7 +198,12 @@ async def reidentify_cache(domain: str, torrent_hash: str,
# 查找指定种子 # 查找指定种子
target_context = None target_context = None
for context in cache_data[domain]: for context in cache_data[domain]:
if HashUtils.md5(f"{context.torrent_info.title}{context.torrent_info.description}") == torrent_hash: if (
HashUtils.md5(
f"{context.torrent_info.title}{context.torrent_info.description}"
)
== torrent_hash
):
target_context = context target_context = context
break break
@@ -167,10 +211,15 @@ async def reidentify_cache(domain: str, torrent_hash: str,
return schemas.Response(success=False, message="未找到指定的种子") return schemas.Response(success=False, message="未找到指定的种子")
# 重新识别 # 重新识别
meta = MetaInfo(title=target_context.torrent_info.title, subtitle=target_context.torrent_info.description) meta = MetaInfo(
title=target_context.torrent_info.title,
subtitle=target_context.torrent_info.description,
)
if tmdbid or doubanid: if tmdbid or doubanid:
# 手动指定媒体信息 # 手动指定媒体信息
mediainfo = await media_chain.async_recognize_media(meta=meta, tmdbid=tmdbid, doubanid=doubanid) mediainfo = await media_chain.async_recognize_media(
meta=meta, tmdbid=tmdbid, doubanid=doubanid
)
else: else:
# 自动重新识别 # 自动重新识别
mediainfo = await media_chain.async_recognize_by_meta(meta) mediainfo = await media_chain.async_recognize_by_meta(meta)
@@ -188,10 +237,16 @@ async def reidentify_cache(domain: str, torrent_hash: str,
# 保存更新后的缓存 # 保存更新后的缓存
await torrents_chain.async_save_cache(cache_data, TorrentsChain().cache_file) await torrents_chain.async_save_cache(cache_data, TorrentsChain().cache_file)
return schemas.Response(success=True, message="重新识别完成", data={ return schemas.Response(
success=True,
message="重新识别完成",
data={
"media_name": mediainfo.title if mediainfo else "", "media_name": mediainfo.title if mediainfo else "",
"media_year": mediainfo.year if mediainfo else "", "media_year": mediainfo.year if mediainfo else "",
"media_type": mediainfo.type.value if mediainfo and mediainfo.type else "" "media_type": mediainfo.type.value
}) if mediainfo and mediainfo.type
else "",
},
)
except Exception as e: except Exception as e:
return schemas.Response(success=False, message=f"重新识别失败:{str(e)}") return schemas.Response(success=False, message=f"重新识别失败:{str(e)}")
+53 -21
View File
@@ -21,8 +21,9 @@ router = APIRouter()
@router.get("/name", summary="查询整理后的名称", response_model=schemas.Response) @router.get("/name", summary="查询整理后的名称", response_model=schemas.Response)
def query_name(path: str, filetype: str, def query_name(
_: schemas.TokenPayload = Depends(verify_token)) -> Any: path: str, filetype: str, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
查询整理后的名称 查询整理后的名称
:param path: 文件路径 :param path: 文件路径
@@ -35,7 +36,9 @@ def query_name(path: str, filetype: str,
) )
if not context or not context.media_info: if not context or not context.media_info:
return schemas.Response(success=False, message="未识别到媒体信息") return schemas.Response(success=False, message="未识别到媒体信息")
new_path = TransferChain().recommend_name(meta=context.meta_info, mediainfo=context.media_info) new_path = TransferChain().recommend_name(
meta=context.meta_info, mediainfo=context.media_info
)
if not new_path: if not new_path:
return schemas.Response(success=False, message="未识别到新名称") return schemas.Response(success=False, message="未识别到新名称")
if filetype == "dir": if filetype == "dir":
@@ -54,9 +57,7 @@ def query_name(path: str, filetype: str,
new_name = parents[0].name new_name = parents[0].name
else: else:
new_name = Path(new_path).name new_name = Path(new_path).name
return schemas.Response(success=True, data={ return schemas.Response(success=True, data={"name": new_name})
"name": new_name
})
@router.get("/queue", summary="查询整理队列", response_model=List[schemas.TransferJob]) @router.get("/queue", summary="查询整理队列", response_model=List[schemas.TransferJob])
@@ -68,8 +69,12 @@ async def query_queue(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
return TransferChain().get_queue_tasks() return TransferChain().get_queue_tasks()
@router.delete("/queue", summary="从整理队列中删除任务", response_model=schemas.Response) @router.delete(
async def remove_queue(fileitem: schemas.FileItem, _: schemas.TokenPayload = Depends(verify_token)) -> Any: "/queue", summary="从整理队列中删除任务", response_model=schemas.Response
)
async def remove_queue(
fileitem: schemas.FileItem, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
查询整理队列 查询整理队列
:param fileitem: 文件项 :param fileitem: 文件项
@@ -82,10 +87,12 @@ async def remove_queue(fileitem: schemas.FileItem, _: schemas.TokenPayload = Dep
@router.post("/manual", summary="手动转移", response_model=schemas.Response) @router.post("/manual", summary="手动转移", response_model=schemas.Response)
def manual_transfer(transer_item: ManualTransferItem, def manual_transfer(
transer_item: ManualTransferItem,
background: Optional[bool] = False, background: Optional[bool] = False,
db: Session = Depends(get_db), db: Session = Depends(get_db),
_: User = Depends(get_current_active_superuser)) -> Any: _: User = Depends(get_current_active_superuser),
) -> Any:
""" """
手动转移文件或历史记录支持自定义剧集识别格式 手动转移文件或历史记录支持自定义剧集识别格式
:param transer_item: 手工整理项 :param transer_item: 手工整理项
@@ -101,7 +108,9 @@ def manual_transfer(transer_item: ManualTransferItem,
# 查询历史记录 # 查询历史记录
history: TransferHistory = TransferHistory.get(db, transer_item.logid) history: TransferHistory = TransferHistory.get(db, transer_item.logid)
if not history: if not history:
return schemas.Response(success=False, message=f"整理记录不存在,ID{transer_item.logid}") return schemas.Response(
success=False, message=f"整理记录不存在,ID{transer_item.logid}"
)
# 强制转移 # 强制转移
force = True force = True
downloader = history.downloader downloader = history.downloader
@@ -118,21 +127,38 @@ def manual_transfer(transer_item: ManualTransferItem,
dest_fileitem = FileItem(**history.dest_fileitem) dest_fileitem = FileItem(**history.dest_fileitem)
state = StorageChain().delete_media_file(dest_fileitem) state = StorageChain().delete_media_file(dest_fileitem)
if not state: if not state:
return schemas.Response(success=False, message=f"{dest_fileitem.path} 删除失败") return schemas.Response(
success=False, message=f"{dest_fileitem.path} 删除失败"
)
# 从历史数据获取信息 # 从历史数据获取信息
if transer_item.from_history: if transer_item.from_history:
transer_item.type_name = history.type if history.type else transer_item.type_name transer_item.type_name = (
transer_item.tmdbid = int(history.tmdbid) if history.tmdbid else transer_item.tmdbid history.type if history.type else transer_item.type_name
transer_item.doubanid = str(history.doubanid) if history.doubanid else transer_item.doubanid )
transer_item.season = int(str(history.seasons).replace("S", "")) if history.seasons else transer_item.season transer_item.tmdbid = (
transer_item.episode_group = history.episode_group or transer_item.episode_group int(history.tmdbid) if history.tmdbid else transer_item.tmdbid
)
transer_item.doubanid = (
str(history.doubanid) if history.doubanid else transer_item.doubanid
)
transer_item.season = (
int(str(history.seasons).replace("S", ""))
if history.seasons
else transer_item.season
)
transer_item.episode_group = (
history.episode_group or transer_item.episode_group
)
if history.episodes: if history.episodes:
if "-" in str(history.episodes): if "-" in str(history.episodes):
# E01-E03多集合并 # E01-E03多集合并
episode_start, episode_end = str(history.episodes).split("-") episode_start, episode_end = str(history.episodes).split("-")
episode_list: list[int] = [] episode_list: list[int] = []
for i in range(int(episode_start.replace("E", "")), int(episode_end.replace("E", "")) + 1): for i in range(
int(episode_start.replace("E", "")),
int(episode_end.replace("E", "")) + 1,
):
episode_list.append(i) episode_list.append(i)
transer_item.episode_detail = ",".join(str(e) for e in episode_list) transer_item.episode_detail = ",".join(str(e) for e in episode_list)
else: else:
@@ -151,11 +177,17 @@ def manual_transfer(transer_item: ManualTransferItem,
try: try:
mtype = MediaType(type_name) mtype = MediaType(type_name)
except ValueError: except ValueError:
return schemas.Response(success=False, message=f"不支持的媒体类型:{type_name}") return schemas.Response(
success=False, message=f"不支持的媒体类型:{type_name}"
)
# 自定义格式 # 自定义格式
epformat = None epformat = None
if transer_item.episode_offset or transer_item.episode_part \ if (
or transer_item.episode_detail or transer_item.episode_format: transer_item.episode_offset
or transer_item.episode_part
or transer_item.episode_detail
or transer_item.episode_format
):
epformat = schemas.EpisodeFormat( epformat = schemas.EpisodeFormat(
format=transer_item.episode_format, format=transer_item.episode_format,
detail=transer_item.episode_detail, detail=transer_item.episode_detail,
+24 -21
View File
@@ -9,8 +9,11 @@ from app import schemas
from app.core.security import get_password_hash from app.core.security import get_password_hash
from app.db import get_async_db from app.db import get_async_db
from app.db.models.user import User from app.db.models.user import User
from app.db.user_oper import get_current_active_superuser_async, \ from app.db.user_oper import (
get_current_active_user_async, get_current_active_user get_current_active_superuser_async,
get_current_active_user_async,
get_current_active_user,
)
from app.db.userconfig_oper import UserConfigOper from app.db.userconfig_oper import UserConfigOper
router = APIRouter() router = APIRouter()
@@ -61,10 +64,12 @@ async def update_user(
user_info = user_in.model_dump() user_info = user_in.model_dump()
if user_info.get("password"): if user_info.get("password"):
# 正则表达式匹配密码包含字母、数字、特殊字符中的至少两项 # 正则表达式匹配密码包含字母、数字、特殊字符中的至少两项
pattern = r'^(?![a-zA-Z]+$)(?!\d+$)(?![^\da-zA-Z\s]+$).{6,50}$' pattern = r"^(?![a-zA-Z]+$)(?!\d+$)(?![^\da-zA-Z\s]+$).{6,50}$"
if not re.match(pattern, user_info.get("password")): if not re.match(pattern, user_info.get("password")):
return schemas.Response(success=False, return schemas.Response(
message="密码需要同时包含字母、数字、特殊字符中的至少两项,且长度大于6位") success=False,
message="密码需要同时包含字母、数字、特殊字符中的至少两项,且长度大于6位",
)
user_info["hashed_password"] = get_password_hash(user_info["password"]) user_info["hashed_password"] = get_password_hash(user_info["password"])
user_info.pop("password") user_info.pop("password")
user = await current_user.async_get_by_id(db, user_id=user_info["id"]) user = await current_user.async_get_by_id(db, user_id=user_info["id"])
@@ -84,7 +89,7 @@ async def update_user(
@router.get("/current", summary="当前登录用户信息", response_model=schemas.User) @router.get("/current", summary="当前登录用户信息", response_model=schemas.User)
async def read_current_user( async def read_current_user(
current_user: User = Depends(get_current_active_user_async) current_user: User = Depends(get_current_active_user_async),
) -> Any: ) -> Any:
""" """
当前登录用户信息 当前登录用户信息
@@ -92,9 +97,15 @@ async def read_current_user(
return current_user return current_user
@router.post("/avatar/{user_id}", summary="上传用户头像", response_model=schemas.Response) @router.post(
async def upload_avatar(user_id: int, db: AsyncSession = Depends(get_async_db), file: UploadFile = File(...), "/avatar/{user_id}", summary="上传用户头像", response_model=schemas.Response
_: User = Depends(get_current_active_user_async)): )
async def upload_avatar(
user_id: int,
db: AsyncSession = Depends(get_async_db),
file: UploadFile = File(...),
_: User = Depends(get_current_active_user_async),
):
""" """
上传用户头像 上传用户头像
""" """
@@ -104,22 +115,17 @@ async def upload_avatar(user_id: int, db: AsyncSession = Depends(get_async_db),
user = await User.async_get(db, user_id) user = await User.async_get(db, user_id)
if not user: if not user:
return schemas.Response(success=False, message="用户不存在") return schemas.Response(success=False, message="用户不存在")
await user.async_update(db, { await user.async_update(db, {"avatar": f"data:image/ico;base64,{file_base64}"})
"avatar": f"data:image/ico;base64,{file_base64}"
})
return schemas.Response(success=True, message=file.filename) return schemas.Response(success=True, message=file.filename)
@router.get("/config/{key}", summary="查询用户配置", response_model=schemas.Response) @router.get("/config/{key}", summary="查询用户配置", response_model=schemas.Response)
def get_config(key: str, def get_config(key: str, current_user: User = Depends(get_current_active_user)):
current_user: User = Depends(get_current_active_user)):
""" """
查询用户配置 查询用户配置
""" """
value = UserConfigOper().get(username=current_user.name, key=key) value = UserConfigOper().get(username=current_user.name, key=key)
return schemas.Response(success=True, data={ return schemas.Response(success=True, data={"value": value})
"value": value
})
@router.post("/config/{key}", summary="更新用户配置", response_model=schemas.Response) @router.post("/config/{key}", summary="更新用户配置", response_model=schemas.Response)
@@ -187,8 +193,5 @@ async def read_user_by_name(
if user == current_user: if user == current_user:
return user return user
if not current_user.is_superuser: if not current_user.is_superuser:
raise HTTPException( raise HTTPException(status_code=400, detail="用户权限不足")
status_code=400,
detail="用户权限不足"
)
return user return user
+8 -4
View File
@@ -17,9 +17,10 @@ def start_webhook_chain(body: Any, form: Any, args: Any):
@router.post("/", summary="Webhook消息响应", response_model=schemas.Response) @router.post("/", summary="Webhook消息响应", response_model=schemas.Response)
async def webhook_message(background_tasks: BackgroundTasks, async def webhook_message(
background_tasks: BackgroundTasks,
request: Request, request: Request,
_: Annotated[str, Depends(verify_apitoken)] _: Annotated[str, Depends(verify_apitoken)],
) -> Any: ) -> Any:
""" """
Webhook响应配置请求中需要添加参数token=API_TOKEN&source=媒体服务器名 Webhook响应配置请求中需要添加参数token=API_TOKEN&source=媒体服务器名
@@ -32,8 +33,11 @@ async def webhook_message(background_tasks: BackgroundTasks,
@router.get("/", summary="Webhook消息响应", response_model=schemas.Response) @router.get("/", summary="Webhook消息响应", response_model=schemas.Response)
async def webhook_message(background_tasks: BackgroundTasks, async def webhook_message_get(
request: Request, _: Annotated[str, Depends(verify_apitoken)]) -> Any: background_tasks: BackgroundTasks,
request: Request,
_: Annotated[str, Depends(verify_apitoken)],
) -> Any:
""" """
Webhook响应配置请求中需要添加参数token=API_TOKEN&source=媒体服务器名 Webhook响应配置请求中需要添加参数token=API_TOKEN&source=媒体服务器名
""" """
+79 -38
View File
@@ -24,8 +24,10 @@ router = APIRouter()
@router.get("/", summary="所有工作流", response_model=List[schemas.Workflow]) @router.get("/", summary="所有工作流", response_model=List[schemas.Workflow])
async def list_workflows(db: AsyncSession = Depends(get_async_db), async def list_workflows(
_: schemas.TokenPayload = Depends(verify_token)) -> Any: db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
获取工作流列表 获取工作流列表
""" """
@@ -33,9 +35,11 @@ async def list_workflows(db: AsyncSession = Depends(get_async_db),
@router.post("/", summary="创建工作流", response_model=schemas.Response) @router.post("/", summary="创建工作流", response_model=schemas.Response)
async def create_workflow(workflow: schemas.Workflow, async def create_workflow(
workflow: schemas.Workflow,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
创建工作流 创建工作流
""" """
@@ -53,7 +57,9 @@ async def create_workflow(workflow: schemas.Workflow,
@router.get("/plugin/actions", summary="查询插件动作", response_model=List[dict]) @router.get("/plugin/actions", summary="查询插件动作", response_model=List[dict])
def list_plugin_actions(plugin_id: str = None, _: schemas.TokenPayload = Depends(verify_token)) -> Any: def list_plugin_actions(
plugin_id: str = None, _: schemas.TokenPayload = Depends(verify_token)
) -> Any:
""" """
获取所有动作 获取所有动作
""" """
@@ -73,33 +79,40 @@ async def get_event_types(_: schemas.TokenPayload = Depends(verify_token)) -> An
""" """
获取所有事件类型 获取所有事件类型
""" """
return [{ return [
{
"title": EVENT_TYPE_NAMES.get(event_type, event_type.name), "title": EVENT_TYPE_NAMES.get(event_type, event_type.name),
"value": event_type.value "value": event_type.value,
} for event_type in EventType] }
for event_type in EventType
]
@router.post("/share", summary="分享工作流", response_model=schemas.Response) @router.post("/share", summary="分享工作流", response_model=schemas.Response)
async def workflow_share( async def workflow_share(
workflow: schemas.WorkflowShare, workflow: schemas.WorkflowShare, _: schemas.TokenPayload = Depends(verify_token)
_: schemas.TokenPayload = Depends(verify_token)) -> Any: ) -> Any:
""" """
分享工作流 分享工作流
""" """
if not workflow.id or not workflow.share_title or not workflow.share_user: if not workflow.id or not workflow.share_title or not workflow.share_user:
return schemas.Response(success=False, message="请填写工作流ID、分享标题和分享人") return schemas.Response(
success=False, message="请填写工作流ID、分享标题和分享人"
)
state, errmsg = await WorkflowHelper().async_workflow_share(workflow_id=workflow.id, state, errmsg = await WorkflowHelper().async_workflow_share(
workflow_id=workflow.id,
share_title=workflow.share_title or "", share_title=workflow.share_title or "",
share_comment=workflow.share_comment or "", share_comment=workflow.share_comment or "",
share_user=workflow.share_user or "") share_user=workflow.share_user or "",
)
return schemas.Response(success=state, message=errmsg) return schemas.Response(success=state, message=errmsg)
@router.delete("/share/{share_id}", summary="删除分享", response_model=schemas.Response) @router.delete("/share/{share_id}", summary="删除分享", response_model=schemas.Response)
async def workflow_share_delete( async def workflow_share_delete(
share_id: int, share_id: int, _: schemas.TokenPayload = Depends(verify_token)
_: schemas.TokenPayload = Depends(verify_token)) -> Any: ) -> Any:
""" """
删除分享 删除分享
""" """
@@ -111,7 +124,8 @@ async def workflow_share_delete(
async def workflow_fork( async def workflow_fork(
workflow: schemas.WorkflowShare, workflow: schemas.WorkflowShare,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.User = Depends(verify_token)) -> Any: _: schemas.User = Depends(verify_token),
) -> Any:
""" """
复用工作流 复用工作流
""" """
@@ -141,11 +155,13 @@ async def workflow_fork(
"timer": workflow.timer, "timer": workflow.timer,
"trigger_type": workflow.trigger_type or "timer", "trigger_type": workflow.trigger_type or "timer",
"event_type": workflow.event_type, "event_type": workflow.event_type,
"event_conditions": json.loads(workflow.event_conditions or "{}") if workflow.event_conditions else {}, "event_conditions": json.loads(workflow.event_conditions or "{}")
if workflow.event_conditions
else {},
"actions": actions, "actions": actions,
"flows": flows, "flows": flows,
"context": context, "context": context,
"state": "P" # 默认暂停状态 "state": "P", # 默认暂停状态
} }
# 检查名称是否重复 # 检查名称是否重复
@@ -163,22 +179,29 @@ async def workflow_fork(
return schemas.Response(success=True, message="复用成功") return schemas.Response(success=True, message="复用成功")
@router.get("/shares", summary="查询分享的工作流", response_model=List[schemas.WorkflowShare]) @router.get(
"/shares", summary="查询分享的工作流", response_model=List[schemas.WorkflowShare]
)
async def workflow_shares( async def workflow_shares(
name: Optional[str] = None, name: Optional[str] = None,
page: Optional[int] = 1, page: Optional[int] = 1,
count: Optional[int] = 30, count: Optional[int] = 30,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
查询分享的工作流 查询分享的工作流
""" """
return await WorkflowHelper().async_get_shares(name=name, page=page, count=count) return await WorkflowHelper().async_get_shares(name=name, page=page, count=count)
@router.post("/{workflow_id}/run", summary="执行工作流", response_model=schemas.Response) @router.post(
def run_workflow(workflow_id: int, "/{workflow_id}/run", summary="执行工作流", response_model=schemas.Response
)
def run_workflow(
workflow_id: int,
from_begin: Optional[bool] = True, from_begin: Optional[bool] = True,
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
执行工作流 执行工作流
""" """
@@ -188,10 +211,14 @@ def run_workflow(workflow_id: int,
return schemas.Response(success=True) return schemas.Response(success=True)
@router.post("/{workflow_id}/start", summary="启用工作流", response_model=schemas.Response) @router.post(
def start_workflow(workflow_id: int, "/{workflow_id}/start", summary="启用工作流", response_model=schemas.Response
)
def start_workflow(
workflow_id: int,
db: Session = Depends(get_db), db: Session = Depends(get_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
启用工作流 启用工作流
""" """
@@ -209,10 +236,14 @@ def start_workflow(workflow_id: int,
return schemas.Response(success=True) return schemas.Response(success=True)
@router.post("/{workflow_id}/pause", summary="停用工作流", response_model=schemas.Response) @router.post(
def pause_workflow(workflow_id: int, "/{workflow_id}/pause", summary="停用工作流", response_model=schemas.Response
)
def pause_workflow(
workflow_id: int,
db: Session = Depends(get_db), db: Session = Depends(get_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
停用工作流 停用工作流
""" """
@@ -233,10 +264,14 @@ def pause_workflow(workflow_id: int,
return schemas.Response(success=True) return schemas.Response(success=True)
@router.post("/{workflow_id}/reset", summary="重置工作流", response_model=schemas.Response) @router.post(
async def reset_workflow(workflow_id: int, "/{workflow_id}/reset", summary="重置工作流", response_model=schemas.Response
)
async def reset_workflow(
workflow_id: int,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
重置工作流 重置工作流
""" """
@@ -253,9 +288,11 @@ async def reset_workflow(workflow_id: int,
@router.get("/{workflow_id}", summary="工作流详情", response_model=schemas.Workflow) @router.get("/{workflow_id}", summary="工作流详情", response_model=schemas.Workflow)
async def get_workflow(workflow_id: int, async def get_workflow(
workflow_id: int,
db: AsyncSession = Depends(get_async_db), db: AsyncSession = Depends(get_async_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
获取工作流详情 获取工作流详情
""" """
@@ -263,9 +300,11 @@ async def get_workflow(workflow_id: int,
@router.put("/{workflow_id}", summary="更新工作流", response_model=schemas.Response) @router.put("/{workflow_id}", summary="更新工作流", response_model=schemas.Response)
def update_workflow(workflow: schemas.Workflow, def update_workflow(
workflow: schemas.Workflow,
db: Session = Depends(get_db), db: Session = Depends(get_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
更新工作流 更新工作流
""" """
@@ -288,9 +327,11 @@ def update_workflow(workflow: schemas.Workflow,
@router.delete("/{workflow_id}", summary="删除工作流", response_model=schemas.Response) @router.delete("/{workflow_id}", summary="删除工作流", response_model=schemas.Response)
def delete_workflow(workflow_id: int, def delete_workflow(
workflow_id: int,
db: Session = Depends(get_db), db: Session = Depends(get_db),
_: schemas.TokenPayload = Depends(verify_token)) -> Any: _: schemas.TokenPayload = Depends(verify_token),
) -> Any:
""" """
删除工作流 删除工作流
""" """
+21 -9
View File
@@ -51,11 +51,15 @@ def extract_text_and_images(content: Any) -> Tuple[str, List[str]]:
data = source.get("data") data = source.get("data")
media_type = source.get("media_type") or "image/png" media_type = source.get("media_type") or "image/png"
if data and str(data).strip(): if data and str(data).strip():
image_urls.append(f"data:{media_type};base64,{str(data).strip()}") image_urls.append(
f"data:{media_type};base64,{str(data).strip()}"
)
return "\n".join(text_parts).strip(), image_urls return "\n".join(text_parts).strip(), image_urls
def build_prompt(messages: List[Any], use_server_session: bool) -> Tuple[str, List[str]]: def build_prompt(
messages: List[Any], use_server_session: bool
) -> Tuple[str, List[str]]:
system_texts: List[str] = [] system_texts: List[str] = []
transcript: List[str] = [] transcript: List[str] = []
latest_user_text = "" latest_user_text = ""
@@ -97,7 +101,9 @@ def build_prompt(messages: List[Any], use_server_session: bool) -> Tuple[str, Li
else: else:
prompt_parts.append("当前用户消息:\n请结合图片内容回复。") prompt_parts.append("当前用户消息:\n请结合图片内容回复。")
return "\n\n".join(part for part in prompt_parts if part).strip(), latest_user_images return "\n\n".join(
part for part in prompt_parts if part
).strip(), latest_user_images
def build_session_id(session_key: str, prefix: str) -> str: def build_session_id(session_key: str, prefix: str) -> str:
@@ -153,18 +159,24 @@ def build_responses_input(
content = item.get("content") content = item.get("content")
messages.append({"role": role, "content": content}) messages.append({"role": role, "content": content})
elif item.get("role") and "content" in item: elif item.get("role") and "content" in item:
messages.append({"role": item.get("role"), "content": item.get("content")}) messages.append(
{"role": item.get("role"), "content": item.get("content")}
)
return messages return messages
if isinstance(input_data, dict) and input_data.get("role") and "content" in input_data: if (
messages.append({"role": input_data.get("role"), "content": input_data.get("content")}) isinstance(input_data, dict)
and input_data.get("role")
and "content" in input_data
):
messages.append(
{"role": input_data.get("role"), "content": input_data.get("content")}
)
return messages return messages
def build_anthropic_messages( def build_anthropic_messages(system: Any, messages: List[Any]) -> List[Dict[str, Any]]:
system: Any, messages: List[Any]
) -> List[Dict[str, Any]]:
normalized: List[Dict[str, Any]] = [] normalized: List[Dict[str, Any]] = []
system_text, _ = extract_text_and_images(system) system_text, _ = extract_text_and_images(system)
if system_text: if system_text:
+141 -127
View File
@@ -16,7 +16,7 @@ from app.schemas import RadarrMovie, SonarrSeries
from app.schemas.types import MediaType from app.schemas.types import MediaType
from version import APP_VERSION from version import APP_VERSION
arr_router = APIRouter(tags=['servarr']) arr_router = APIRouter(tags=["servarr"])
@arr_router.get("/system/status", summary="系统状态") @arr_router.get("/system/status", summary="系统状态")
@@ -51,7 +51,7 @@ async def arr_system_status(_: Annotated[str, Depends(verify_apikey)]) -> Any:
"build": 0, "build": 0,
"revision": 0, "revision": 0,
"majorRevision": 0, "majorRevision": 0,
"minorRevision": 0 "minorRevision": 0,
}, },
"authentication": "none", "authentication": "none",
"migrationVersion": 0, "migrationVersion": 0,
@@ -62,14 +62,14 @@ async def arr_system_status(_: Annotated[str, Depends(verify_apikey)]) -> Any:
"build": 0, "build": 0,
"revision": 0, "revision": 0,
"majorRevision": 0, "majorRevision": 0,
"minorRevision": 0 "minorRevision": 0,
}, },
"runtimeName": "", "runtimeName": "",
"startTime": "", "startTime": "",
"packageVersion": "", "packageVersion": "",
"packageAuthor": "jxxghp", "packageAuthor": "jxxghp",
"packageUpdateMechanism": "builtIn", "packageUpdateMechanism": "builtIn",
"packageUpdateMechanismMessage": "" "packageUpdateMechanismMessage": "",
} }
@@ -92,24 +92,15 @@ async def arr_qualityProfile(_: Annotated[str, Depends(verify_apikey)]) -> Any:
"id": 0, "id": 0,
"name": "默认", "name": "默认",
"source": "0", "source": "0",
"resolution": 0 "resolution": 0,
}, },
"items": [ "items": ["string"],
"string" "allowed": True,
],
"allowed": True
} }
], ],
"minFormatScore": 0, "minFormatScore": 0,
"cutoffFormatScore": 0, "cutoffFormatScore": 0,
"formatItems": [ "formatItems": [{"id": 0, "format": 0, "name": "默认", "score": 0}],
{
"id": 0,
"format": 0,
"name": "默认",
"score": 0
}
]
} }
] ]
@@ -125,7 +116,7 @@ async def arr_rootfolder(_: Annotated[str, Depends(verify_apikey)]) -> Any:
"path": "/", "path": "/",
"accessible": True, "accessible": True,
"freeSpace": 0, "freeSpace": 0,
"unmappedFolders": [] "unmappedFolders": [],
} }
] ]
@@ -135,12 +126,7 @@ async def arr_tag(_: Annotated[str, Depends(verify_apikey)]) -> Any:
""" """
模拟RadarrSonarr标签 模拟RadarrSonarr标签
""" """
return [ return [{"id": 1, "label": "默认"}]
{
"id": 1,
"label": "默认"
}
]
@arr_router.get("/languageprofile", summary="语言") @arr_router.get("/languageprofile", summary="语言")
@@ -148,29 +134,25 @@ async def arr_languageprofile(_: Annotated[str, Depends(verify_apikey)]) -> Any:
""" """
模拟RadarrSonarr语言 模拟RadarrSonarr语言
""" """
return [{ return [
{
"id": 1, "id": 1,
"name": "默认", "name": "默认",
"upgradeAllowed": True, "upgradeAllowed": True,
"cutoff": { "cutoff": {"id": 1, "name": "默认"},
"id": 1,
"name": "默认"
},
"languages": [ "languages": [
{ {"id": 1, "language": {"id": 1, "name": "默认"}, "allowed": True}
"id": 1, ],
"language": {
"id": 1,
"name": "默认"
},
"allowed": True
} }
] ]
}]
@arr_router.get("/movie", summary="所有订阅电影", response_model=List[schemas.RadarrMovie]) @arr_router.get(
async def arr_movies(_: Annotated[str, Depends(verify_apikey)], db: AsyncSession = Depends(get_async_db)) -> Any: "/movie", summary="所有订阅电影", response_model=List[schemas.RadarrMovie]
)
async def arr_movies(
_: Annotated[str, Depends(verify_apikey)], db: AsyncSession = Depends(get_async_db)
) -> Any:
""" """
查询Rardar电影 查询Rardar电影
""" """
@@ -245,7 +227,8 @@ async def arr_movies(_: Annotated[str, Depends(verify_apikey)], db: AsyncSession
for subscribe in subscribes: for subscribe in subscribes:
if subscribe.type != MediaType.MOVIE.value: if subscribe.type != MediaType.MOVIE.value:
continue continue
result.append(RadarrMovie( result.append(
RadarrMovie(
id=subscribe.id, id=subscribe.id,
title=subscribe.name, title=subscribe.name,
year=subscribe.year, year=subscribe.year,
@@ -255,13 +238,18 @@ async def arr_movies(_: Annotated[str, Depends(verify_apikey)], db: AsyncSession
imdbId=subscribe.imdbid, imdbId=subscribe.imdbid,
profileId=1, profileId=1,
qualityProfileId=1, qualityProfileId=1,
hasFile=False hasFile=False,
)) )
)
return result return result
@arr_router.get("/movie/lookup", summary="查询电影", response_model=List[schemas.RadarrMovie]) @arr_router.get(
def arr_movie_lookup(term: str, _: Annotated[str, Depends(verify_apikey)], db: Session = Depends(get_db)) -> Any: "/movie/lookup", summary="查询电影", response_model=List[schemas.RadarrMovie]
)
def arr_movie_lookup(
term: str, _: Annotated[str, Depends(verify_apikey)], db: Session = Depends(get_db)
) -> Any:
""" """
查询Rardar电影 term: `tmdb:${id}` 查询Rardar电影 term: `tmdb:${id}`
存在和不存在均不能返回错误 存在和不存在均不能返回错误
@@ -290,7 +278,8 @@ def arr_movie_lookup(term: str, _: Annotated[str, Depends(verify_apikey)], db: S
subid = None subid = None
monitored = False monitored = False
return [RadarrMovie( return [
RadarrMovie(
id=subid, id=subid,
title=mediainfo.title, title=mediainfo.title,
year=mediainfo.year, year=mediainfo.year,
@@ -302,13 +291,19 @@ def arr_movie_lookup(term: str, _: Annotated[str, Depends(verify_apikey)], db: S
folderName=mediainfo.title_year, folderName=mediainfo.title_year,
profileId=1, profileId=1,
qualityProfileId=1, qualityProfileId=1,
hasFile=hasfile hasFile=hasfile,
)] )
]
@arr_router.get("/movie/{mid}", summary="电影订阅详情", response_model=schemas.RadarrMovie) @arr_router.get(
async def arr_movie(mid: int, _: Annotated[str, Depends(verify_apikey)], "/movie/{mid}", summary="电影订阅详情", response_model=schemas.RadarrMovie
db: AsyncSession = Depends(get_async_db)) -> Any: )
async def arr_movie(
mid: int,
_: Annotated[str, Depends(verify_apikey)],
db: AsyncSession = Depends(get_async_db),
) -> Any:
""" """
查询Rardar电影订阅 查询Rardar电影订阅
""" """
@@ -324,19 +319,17 @@ async def arr_movie(mid: int, _: Annotated[str, Depends(verify_apikey)],
imdbId=subscribe.imdbid, imdbId=subscribe.imdbid,
profileId=1, profileId=1,
qualityProfileId=1, qualityProfileId=1,
hasFile=False hasFile=False,
) )
else: else:
raise HTTPException( raise HTTPException(status_code=404, detail="未找到该电影!")
status_code=404,
detail="未找到该电影!"
)
@arr_router.post("/movie", summary="新增电影订阅") @arr_router.post("/movie", summary="新增电影订阅")
async def arr_add_movie(_: Annotated[str, Depends(verify_apikey)], async def arr_add_movie(
_: Annotated[str, Depends(verify_apikey)],
movie: RadarrMovie, movie: RadarrMovie,
db: AsyncSession = Depends(get_async_db) db: AsyncSession = Depends(get_async_db),
) -> Any: ) -> Any:
""" """
新增Rardar电影订阅 新增Rardar电影订阅
@@ -344,29 +337,29 @@ async def arr_add_movie(_: Annotated[str, Depends(verify_apikey)],
# 检查订阅是否已存在 # 检查订阅是否已存在
subscribe = await Subscribe.async_get_by_tmdbid(db, movie.tmdbId) subscribe = await Subscribe.async_get_by_tmdbid(db, movie.tmdbId)
if subscribe: if subscribe:
return { return {"id": subscribe.id}
"id": subscribe.id
}
# 添加订阅 # 添加订阅
sid, message = await SubscribeChain().async_add(title=movie.title, sid, message = await SubscribeChain().async_add(
title=movie.title,
year=movie.year, year=movie.year,
mtype=MediaType.MOVIE, mtype=MediaType.MOVIE,
tmdbid=movie.tmdbId, tmdbid=movie.tmdbId,
username="Seerr") username="Seerr",
if sid:
return {
"id": sid
}
else:
raise HTTPException(
status_code=500,
detail=f"添加订阅失败:{message}"
) )
if sid:
return {"id": sid}
else:
raise HTTPException(status_code=500, detail=f"添加订阅失败:{message}")
@arr_router.delete("/movie/{mid}", summary="删除电影订阅", response_model=schemas.Response) @arr_router.delete(
async def arr_remove_movie(mid: int, _: Annotated[str, Depends(verify_apikey)], "/movie/{mid}", summary="删除电影订阅", response_model=schemas.Response
db: AsyncSession = Depends(get_async_db)) -> Any: )
async def arr_remove_movie(
mid: int,
_: Annotated[str, Depends(verify_apikey)],
db: AsyncSession = Depends(get_async_db),
) -> Any:
""" """
删除Rardar电影订阅 删除Rardar电影订阅
""" """
@@ -375,14 +368,15 @@ async def arr_remove_movie(mid: int, _: Annotated[str, Depends(verify_apikey)],
await subscribe.async_delete(db, mid) await subscribe.async_delete(db, mid)
return schemas.Response(success=True) return schemas.Response(success=True)
else: else:
raise HTTPException( raise HTTPException(status_code=404, detail="未找到该电影!")
status_code=404,
detail="未找到该电影!"
@arr_router.get(
"/series", summary="所有剧集", response_model=List[schemas.SonarrSeries]
) )
async def arr_series(
_: Annotated[str, Depends(verify_apikey)], db: AsyncSession = Depends(get_async_db)
@arr_router.get("/series", summary="所有剧集", response_model=List[schemas.SonarrSeries]) ) -> Any:
async def arr_series(_: Annotated[str, Depends(verify_apikey)], db: AsyncSession = Depends(get_async_db)) -> Any:
""" """
查询Sonarr剧集 查询Sonarr剧集
""" """
@@ -494,14 +488,17 @@ async def arr_series(_: Annotated[str, Depends(verify_apikey)], db: AsyncSession
for subscribe in subscribes: for subscribe in subscribes:
if subscribe.type != MediaType.TV.value: if subscribe.type != MediaType.TV.value:
continue continue
result.append(SonarrSeries( result.append(
SonarrSeries(
id=subscribe.id, id=subscribe.id,
title=subscribe.name, title=subscribe.name,
seasonCount=1, seasonCount=1,
seasons=[{ seasons=[
{
"seasonNumber": subscribe.season, "seasonNumber": subscribe.season,
"monitored": True, "monitored": True,
}], }
],
remotePoster=subscribe.poster, remotePoster=subscribe.poster,
year=subscribe.year, year=subscribe.year,
tmdbId=subscribe.tmdbid, tmdbId=subscribe.tmdbid,
@@ -512,13 +509,16 @@ async def arr_series(_: Annotated[str, Depends(verify_apikey)], db: AsyncSession
qualityProfileId=1, qualityProfileId=1,
isAvailable=True, isAvailable=True,
monitored=True, monitored=True,
hasFile=False hasFile=False,
)) )
)
return result return result
@arr_router.get("/series/lookup", summary="查询剧集") @arr_router.get("/series/lookup", summary="查询剧集")
def arr_series_lookup(term: str, _: Annotated[str, Depends(verify_apikey)], db: Session = Depends(get_db)) -> Any: def arr_series_lookup(
term: str, _: Annotated[str, Depends(verify_apikey)], db: Session = Depends(get_db)
) -> Any:
""" """
查询Sonarr剧集 term: `tvdb:${id}` title 查询Sonarr剧集 term: `tvdb:${id}` title
""" """
@@ -542,13 +542,19 @@ def arr_series_lookup(term: str, _: Annotated[str, Depends(verify_apikey)], db:
continue continue
# 季信息(只取默认季类型,排除特别季) # 季信息(只取默认季类型,排除特别季)
sea_num = len([season for season in tvdbinfo.get('seasons') if sea_num = len(
season['type']['id'] == tvdbinfo.get('defaultSeasonType') and season['number'] > 0]) [
season
for season in tvdbinfo.get("seasons")
if season["type"]["id"] == tvdbinfo.get("defaultSeasonType")
and season["number"] > 0
]
)
if sea_num: if sea_num:
seas = list(range(1, int(sea_num) + 1)) seas = list(range(1, int(sea_num) + 1))
# 根据TVDB查询媒体信息 # 根据TVDB查询媒体信息
meta = MetaInfo(tvdbinfo.get('name')) meta = MetaInfo(tvdbinfo.get("name"))
meta.type = MediaType.TV meta.type = MediaType.TV
mediainfo = MediaChain().recognize_by_meta( mediainfo = MediaChain().recognize_by_meta(
meta, meta,
@@ -573,24 +579,30 @@ def arr_series_lookup(term: str, _: Annotated[str, Depends(verify_apikey)], db:
sub_seas = [sub.season for sub in subscribes] sub_seas = [sub.season for sub in subscribes]
for sea in seas: for sea in seas:
if sea in sub_seas: if sea in sub_seas:
seasons.append({ seasons.append(
{
"seasonNumber": sea, "seasonNumber": sea,
"monitored": True, "monitored": True,
}) }
)
else: else:
seasons.append({ seasons.append(
{
"seasonNumber": sea, "seasonNumber": sea,
"monitored": False, "monitored": False,
}) }
)
subid = subscribes[-1].id subid = subscribes[-1].id
else: else:
subid = None subid = None
monitored = False monitored = False
for sea in seas: for sea in seas:
seasons.append({ seasons.append(
{
"seasonNumber": sea, "seasonNumber": sea,
"monitored": False, "monitored": False,
}) }
)
sonarr_series = SonarrSeries( sonarr_series = SonarrSeries(
id=subid, id=subid,
title=mediainfo.title, title=mediainfo.title,
@@ -612,8 +624,11 @@ def arr_series_lookup(term: str, _: Annotated[str, Depends(verify_apikey)], db:
@arr_router.get("/series/{tid}", summary="剧集详情") @arr_router.get("/series/{tid}", summary="剧集详情")
async def arr_serie(tid: int, _: Annotated[str, Depends(verify_apikey)], async def arr_serie(
db: AsyncSession = Depends(get_async_db)) -> Any: tid: int,
_: Annotated[str, Depends(verify_apikey)],
db: AsyncSession = Depends(get_async_db),
) -> Any:
""" """
查询Sonarr剧集 查询Sonarr剧集
""" """
@@ -623,10 +638,12 @@ async def arr_serie(tid: int, _: Annotated[str, Depends(verify_apikey)],
id=subscribe.id, id=subscribe.id,
title=subscribe.name, title=subscribe.name,
seasonCount=1, seasonCount=1,
seasons=[{ seasons=[
{
"seasonNumber": subscribe.season, "seasonNumber": subscribe.season,
"monitored": True, "monitored": True,
}], }
],
year=subscribe.year, year=subscribe.year,
remotePoster=subscribe.poster, remotePoster=subscribe.poster,
tmdbId=subscribe.tmdbid, tmdbId=subscribe.tmdbid,
@@ -637,61 +654,58 @@ async def arr_serie(tid: int, _: Annotated[str, Depends(verify_apikey)],
qualityProfileId=1, qualityProfileId=1,
isAvailable=True, isAvailable=True,
monitored=True, monitored=True,
hasFile=False hasFile=False,
) )
else: else:
raise HTTPException( raise HTTPException(status_code=404, detail="未找到该电视剧!")
status_code=404,
detail="未找到该电视剧!"
)
@arr_router.post("/series", summary="新增剧集订阅") @arr_router.post("/series", summary="新增剧集订阅")
async def arr_add_series(tv: schemas.SonarrSeries, async def arr_add_series(
tv: schemas.SonarrSeries,
_: Annotated[str, Depends(verify_apikey)], _: Annotated[str, Depends(verify_apikey)],
db: AsyncSession = Depends(get_async_db)) -> Any: db: AsyncSession = Depends(get_async_db),
) -> Any:
""" """
新增Sonarr剧集订阅 新增Sonarr剧集订阅
""" """
# 检查订阅是否存在 # 检查订阅是否存在
left_seasons = [] left_seasons = []
for season in tv.seasons: for season in tv.seasons:
subscribe = await Subscribe.async_get_by_tmdbid(db, tmdbid=tv.tmdbId, subscribe = await Subscribe.async_get_by_tmdbid(
season=season.get("seasonNumber")) db, tmdbid=tv.tmdbId, season=season.get("seasonNumber")
)
if subscribe: if subscribe:
continue continue
left_seasons.append(season) left_seasons.append(season)
# 全部已存在订阅 # 全部已存在订阅
if not left_seasons: if not left_seasons:
return { return {"id": 1}
"id": 1
}
# 剩下的添加订阅 # 剩下的添加订阅
sid = 0 sid = 0
message = "" message = ""
for season in left_seasons: for season in left_seasons:
if not season.get("monitored"): if not season.get("monitored"):
continue continue
sid, message = await SubscribeChain().async_add(title=tv.title, sid, message = await SubscribeChain().async_add(
title=tv.title,
year=tv.year, year=tv.year,
season=season.get("seasonNumber"), season=season.get("seasonNumber"),
tmdbid=tv.tmdbId, tmdbid=tv.tmdbId,
mtype=MediaType.TV, mtype=MediaType.TV,
username="Seerr") username="Seerr",
)
if sid: if sid:
return { return {"id": sid}
"id": sid
}
else: else:
raise HTTPException( raise HTTPException(status_code=500, detail=f"添加订阅失败:{message}")
status_code=500,
detail=f"添加订阅失败:{message}"
)
@arr_router.put("/series", summary="更新剧集订阅") @arr_router.put("/series", summary="更新剧集订阅")
async def arr_update_series(tv: schemas.SonarrSeries, _: Annotated[str, Depends(verify_apikey)]) -> Any: async def arr_update_series(
tv: schemas.SonarrSeries, _: Annotated[str, Depends(verify_apikey)]
) -> Any:
""" """
更新Sonarr剧集订阅 更新Sonarr剧集订阅
""" """
@@ -699,8 +713,11 @@ async def arr_update_series(tv: schemas.SonarrSeries, _: Annotated[str, Depends(
@arr_router.delete("/series/{tid}", summary="删除剧集订阅") @arr_router.delete("/series/{tid}", summary="删除剧集订阅")
async def arr_remove_series(tid: int, _: Annotated[str, Depends(verify_apikey)], async def arr_remove_series(
db: AsyncSession = Depends(get_async_db)) -> Any: tid: int,
_: Annotated[str, Depends(verify_apikey)],
db: AsyncSession = Depends(get_async_db),
) -> Any:
""" """
删除Sonarr剧集订阅 删除Sonarr剧集订阅
""" """
@@ -709,7 +726,4 @@ async def arr_remove_series(tid: int, _: Annotated[str, Depends(verify_apikey)],
await subscribe.async_delete(db, tid) await subscribe.async_delete(db, tid)
return schemas.Response(success=True) return schemas.Response(success=True)
else: else:
raise HTTPException( raise HTTPException(status_code=404, detail="未找到该电视剧!")
status_code=404,
detail="未找到该电视剧!"
)
+11 -8
View File
@@ -15,7 +15,6 @@ from app.utils.crypto import CryptoJsUtils, HashUtils
class GzipRequest(Request): class GzipRequest(Request):
async def body(self) -> bytes: async def body(self) -> bytes:
if not hasattr(self, "_body"): if not hasattr(self, "_body"):
body = await super().body() body = await super().body()
@@ -26,7 +25,6 @@ class GzipRequest(Request):
class GzipRoute(APIRoute): class GzipRoute(APIRoute):
def get_route_handler(self) -> Callable: def get_route_handler(self) -> Callable:
original_route_handler = super().get_route_handler() original_route_handler = super().get_route_handler()
@@ -46,9 +44,11 @@ async def verify_server_enabled():
return True return True
cookie_router = APIRouter(route_class=GzipRoute, cookie_router = APIRouter(
route_class=GzipRoute,
tags=["servcookie"], tags=["servcookie"],
dependencies=[Depends(verify_server_enabled)]) dependencies=[Depends(verify_server_enabled)],
)
@cookie_router.get("/", response_class=PlainTextResponse) @cookie_router.get("/", response_class=PlainTextResponse)
@@ -95,8 +95,9 @@ async def load_encrypt_data(uuid: str) -> Dict[str, Any]:
return data return data
def get_decrypted_cookie_data(uuid: str, password: str, def get_decrypted_cookie_data(
encrypted: str) -> Optional[Dict[str, Any]]: uuid: str, password: str, encrypted: str
) -> Optional[Dict[str, Any]]:
""" """
加载本地加密数据并解密为Cookie 加载本地加密数据并解密为Cookie
""" """
@@ -118,7 +119,8 @@ def get_decrypted_cookie_data(uuid: str, password: str,
@cookie_router.get("/get/{uuid}") @cookie_router.get("/get/{uuid}")
async def get_cookie( async def get_cookie(
uuid: Annotated[str, Path(min_length=5, pattern="^[a-zA-Z0-9]+$")]): uuid: Annotated[str, Path(min_length=5, pattern="^[a-zA-Z0-9]+$")],
):
""" """
GET 下载加密数据 GET 下载加密数据
""" """
@@ -128,7 +130,8 @@ async def get_cookie(
@cookie_router.post("/get/{uuid}") @cookie_router.post("/get/{uuid}")
async def post_cookie( async def post_cookie(
uuid: Annotated[str, Path(min_length=5, pattern="^[a-zA-Z0-9]+$")], uuid: Annotated[str, Path(min_length=5, pattern="^[a-zA-Z0-9]+$")],
request: Optional[schemas.CookiePassword] = Body(None)): request: Optional[schemas.CookiePassword] = Body(None),
):
""" """
POST 下载加密数据 POST 下载加密数据
""" """
@@ -1,10 +1,5 @@
from ..tmdb import TMDb from ..tmdb import TMDb
try:
from urllib import quote
except ImportError:
from urllib.parse import quote
class TV(TMDb): class TV(TMDb):
_urls = { _urls = {
+227 -72
View File
@@ -3,7 +3,6 @@ import asyncio
import json import json
import tempfile import tempfile
import unittest import unittest
from pathlib import Path
from types import ModuleType, SimpleNamespace from types import ModuleType, SimpleNamespace
from unittest.mock import ANY, MagicMock, patch from unittest.mock import ANY, MagicMock, patch
@@ -21,15 +20,20 @@ if "Pinyin2Hanzi" not in sys.modules:
from app.modules.feishu import FeishuModule from app.modules.feishu import FeishuModule
from app.modules.feishu.feishu import Feishu from app.modules.feishu.feishu import Feishu
from app.schemas import Notification from app.schemas import Notification
from app.schemas.message import ChannelCapability, ChannelCapabilityManager, MessageResponse from app.schemas.message import (
ChannelCapability,
ChannelCapabilityManager,
MessageResponse,
)
from app.schemas.types import MessageChannel, NotificationType from app.schemas.types import MessageChannel, NotificationType
class TestFeishu(unittest.TestCase): class TestFeishu(unittest.TestCase):
@staticmethod @staticmethod
def _build_client(**kwargs) -> Feishu: def _build_client(**kwargs) -> Feishu:
with patch.object(Feishu, "_build_api_client", return_value=MagicMock()), patch.object( with (
Feishu, "_start_ws_client" patch.object(Feishu, "_build_api_client", return_value=MagicMock()),
patch.object(Feishu, "_start_ws_client"),
): ):
return Feishu( return Feishu(
FEISHU_APP_ID="cli_test_app_id", FEISHU_APP_ID="cli_test_app_id",
@@ -64,7 +68,21 @@ class TestFeishu(unittest.TestCase):
return response return response
@staticmethod @staticmethod
def _build_message_api(create_response=None, patch_response=None, reply_response=None, reaction_create_response=None, reaction_delete_response=None, card_create_response=None, card_settings_response=None, card_content_response=None, image_create_response=None, file_create_response=None, image_get_response=None, file_get_response=None, message_resource_response=None): def _build_message_api(
create_response=None,
patch_response=None,
reply_response=None,
reaction_create_response=None,
reaction_delete_response=None,
card_create_response=None,
card_settings_response=None,
card_content_response=None,
image_create_response=None,
file_create_response=None,
image_get_response=None,
file_get_response=None,
message_resource_response=None,
):
message_api = SimpleNamespace( message_api = SimpleNamespace(
create=MagicMock(return_value=create_response), create=MagicMock(return_value=create_response),
patch=MagicMock(return_value=patch_response), patch=MagicMock(return_value=patch_response),
@@ -111,7 +129,11 @@ class TestFeishu(unittest.TestCase):
return api_client, message_api return api_client, message_api
@staticmethod @staticmethod
def _resource_response(content: bytes, file_name: str = "resource.bin", content_type: str = "application/octet-stream"): def _resource_response(
content: bytes,
file_name: str = "resource.bin",
content_type: str = "application/octet-stream",
):
response = MagicMock() response = MagicMock()
response.code = 0 response.code = 0
response.file = MagicMock() response.file = MagicMock()
@@ -169,9 +191,11 @@ class TestFeishu(unittest.TestCase):
class _Builder: class _Builder:
def __getattr__(self, name): def __getattr__(self, name):
if name.startswith("register_"): if name.startswith("register_"):
def _register(handler): def _register(handler):
registered.append(name) registered.append(name)
return self return self
return _register return _register
raise AttributeError(name) raise AttributeError(name)
@@ -181,7 +205,10 @@ class TestFeishu(unittest.TestCase):
client = self._build_client() client = self._build_client()
fake_builder = _Builder() fake_builder = _Builder()
with patch("app.modules.feishu.feishu.lark.EventDispatcherHandler.builder", return_value=fake_builder): with patch(
"app.modules.feishu.feishu.lark.EventDispatcherHandler.builder",
return_value=fake_builder,
):
handler = client._build_event_handler() handler = client._build_event_handler()
self.assertEqual(handler, "handler") self.assertEqual(handler, "handler")
@@ -190,15 +217,20 @@ class TestFeishu(unittest.TestCase):
self.assertIn("register_p2_im_message_reaction_created_v1", registered) self.assertIn("register_p2_im_message_reaction_created_v1", registered)
self.assertIn("register_p2_im_message_reaction_deleted_v1", registered) self.assertIn("register_p2_im_message_reaction_deleted_v1", registered)
self.assertIn("register_p2_im_message_recalled_v1", registered) self.assertIn("register_p2_im_message_recalled_v1", registered)
self.assertIn("register_p2_im_chat_access_event_bot_p2p_chat_entered_v1", registered) self.assertIn(
"register_p2_im_chat_access_event_bot_p2p_chat_entered_v1", registered
)
self.assertIn("register_p2_card_action_trigger", registered) self.assertIn("register_p2_card_action_trigger", registered)
def test_parse_message_blocks_non_admin_command(self): def test_parse_message_blocks_non_admin_command(self):
client = self._build_client(FEISHU_ADMINS="ou_admin") client = self._build_client(FEISHU_ADMINS="ou_admin")
with patch("app.modules.feishu.feishu.UserOper.get_name", return_value=None), patch.object( with (
patch("app.modules.feishu.feishu.UserOper.get_name", return_value=None),
patch.object(
client, "send_text", return_value={"success": True} client, "send_text", return_value={"success": True}
) as send_text: ) as send_text,
):
result = client.parse_message( result = client.parse_message(
{ {
"type": "message", "type": "message",
@@ -223,7 +255,10 @@ class TestFeishu(unittest.TestCase):
def test_parse_message_maps_feishu_ids_to_moviepilot_username(self): def test_parse_message_maps_feishu_ids_to_moviepilot_username(self):
client = self._build_client() client = self._build_client()
with patch("app.modules.feishu.feishu.UserOper.get_name", return_value="moviepilot-user") as get_name: with patch(
"app.modules.feishu.feishu.UserOper.get_name",
return_value="moviepilot-user",
) as get_name:
result = client.parse_message( result = client.parse_message(
{ {
"type": "message", "type": "message",
@@ -317,7 +352,9 @@ class TestFeishu(unittest.TestCase):
self.assertEqual(image_element["img_key"], "img_v2_remote") self.assertEqual(image_element["img_key"], "img_v2_remote")
self.assertEqual(content["body"]["elements"][1]["margin"], "12px 12px 0px 12px") self.assertEqual(content["body"]["elements"][1]["margin"], "12px 12px 0px 12px")
self.assertEqual(content["body"]["elements"][2]["margin"], "4px 12px 12px 12px") self.assertEqual(content["body"]["elements"][2]["margin"], "4px 12px 12px 12px")
self.assertEqual(content["body"]["elements"][-1]["margin"], "0px 12px 12px 12px") self.assertEqual(
content["body"]["elements"][-1]["margin"], "0px 12px 12px 12px"
)
self.assertEqual(content["body"]["elements"][-1]["tag"], "column_set") self.assertEqual(content["body"]["elements"][-1]["tag"], "column_set")
def test_send_notification_supports_user_id_target(self): def test_send_notification_supports_user_id_target(self):
@@ -390,25 +427,33 @@ class TestFeishu(unittest.TestCase):
reaction_delete_response=self._success_response(), reaction_delete_response=self._success_response(),
) )
reaction_id = client.add_message_reaction("om_origin", Feishu.PROCESSING_REACTION_EMOJI) reaction_id = client.add_message_reaction(
"om_origin", Feishu.PROCESSING_REACTION_EMOJI
)
deleted = client.delete_message_reaction("om_origin", "reaction_1") deleted = client.delete_message_reaction("om_origin", "reaction_1")
self.assertEqual(reaction_id, "reaction_1") self.assertEqual(reaction_id, "reaction_1")
self.assertTrue(deleted) self.assertTrue(deleted)
create_request = client._api_client.im.v1.message_reaction.create.call_args.args[0] create_request = (
client._api_client.im.v1.message_reaction.create.call_args.args[0]
)
self.assertEqual(create_request.message_id, "om_origin") self.assertEqual(create_request.message_id, "om_origin")
self.assertEqual( self.assertEqual(
create_request.request_body.reaction_type.emoji_type, create_request.request_body.reaction_type.emoji_type,
Feishu.PROCESSING_REACTION_EMOJI, Feishu.PROCESSING_REACTION_EMOJI,
) )
delete_request = client._api_client.im.v1.message_reaction.delete.call_args.args[0] delete_request = (
client._api_client.im.v1.message_reaction.delete.call_args.args[0]
)
self.assertEqual(delete_request.message_id, "om_origin") self.assertEqual(delete_request.message_id, "om_origin")
self.assertEqual(delete_request.reaction_id, "reaction_1") self.assertEqual(delete_request.reaction_id, "reaction_1")
def test_send_notification_uses_streaming_card_for_agent_text(self): def test_send_notification_uses_streaming_card_for_agent_text(self):
client = self._build_client() client = self._build_client()
client._api_client, message_api = self._build_message_api( client._api_client, message_api = self._build_message_api(
create_response=self._success_response(message_id="om_stream", chat_id="oc_stream"), create_response=self._success_response(
message_id="om_stream", chat_id="oc_stream"
),
card_create_response=self._card_create_success_response("card_stream"), card_create_response=self._card_create_success_response("card_stream"),
) )
@@ -422,21 +467,31 @@ class TestFeishu(unittest.TestCase):
) )
self.assertTrue(result["success"]) self.assertTrue(result["success"])
self.assertEqual(result["metadata"]["feishu_streaming"]["card_id"], "card_stream") self.assertEqual(
result["metadata"]["feishu_streaming"]["card_id"], "card_stream"
)
self.assertEqual(result["metadata"]["feishu_streaming"]["sequence"], 0) self.assertEqual(result["metadata"]["feishu_streaming"]["sequence"], 0)
card_request = client._api_client.cardkit.v1.card.create.call_args.args[0] card_request = client._api_client.cardkit.v1.card.create.call_args.args[0]
self.assertEqual(card_request.request_body.type, "card_json") self.assertEqual(card_request.request_body.type, "card_json")
card_payload = json.loads(card_request.request_body.data) card_payload = json.loads(card_request.request_body.data)
self.assertTrue(card_payload["config"]["streaming_mode"]) self.assertTrue(card_payload["config"]["streaming_mode"])
self.assertEqual(card_payload["body"]["elements"][-1]["element_id"], Feishu.STREAM_CARD_BODY_ELEMENT_ID) self.assertEqual(
card_payload["body"]["elements"][-1]["element_id"],
Feishu.STREAM_CARD_BODY_ELEMENT_ID,
)
message_request = message_api.create.call_args.args[0] message_request = message_api.create.call_args.args[0]
self.assertEqual(message_request.request_body.msg_type, "interactive") self.assertEqual(message_request.request_body.msg_type, "interactive")
self.assertEqual(json.loads(message_request.request_body.content)["data"]["card_id"], "card_stream") self.assertEqual(
json.loads(message_request.request_body.content)["data"]["card_id"],
"card_stream",
)
def test_send_notification_replies_with_streaming_card_for_agent_text(self): def test_send_notification_replies_with_streaming_card_for_agent_text(self):
client = self._build_client() client = self._build_client()
client._api_client, message_api = self._build_message_api( client._api_client, message_api = self._build_message_api(
reply_response=self._success_response(message_id="om_reply", chat_id="oc_stream"), reply_response=self._success_response(
message_id="om_reply", chat_id="oc_stream"
),
card_create_response=self._card_create_success_response("card_stream"), card_create_response=self._card_create_success_response("card_stream"),
) )
@@ -455,7 +510,10 @@ class TestFeishu(unittest.TestCase):
reply_request = message_api.reply.call_args.args[0] reply_request = message_api.reply.call_args.args[0]
self.assertEqual(reply_request.message_id, "om_origin") self.assertEqual(reply_request.message_id, "om_origin")
self.assertEqual(reply_request.request_body.msg_type, "interactive") self.assertEqual(reply_request.request_body.msg_type, "interactive")
self.assertEqual(json.loads(reply_request.request_body.content)["data"]["card_id"], "card_stream") self.assertEqual(
json.loads(reply_request.request_body.content)["data"]["card_id"],
"card_stream",
)
self.assertEqual(result["metadata"]["feishu_streaming"]["sequence"], 0) self.assertEqual(result["metadata"]["feishu_streaming"]["sequence"], 0)
def test_edit_replied_streaming_card_uses_first_increment_sequence(self): def test_edit_replied_streaming_card_uses_first_increment_sequence(self):
@@ -479,7 +537,9 @@ class TestFeishu(unittest.TestCase):
self.assertTrue(success) self.assertTrue(success)
message_api.patch.assert_not_called() message_api.patch.assert_not_called()
content_request = client._api_client.cardkit.v1.card_element.content.call_args.args[0] content_request = (
client._api_client.cardkit.v1.card_element.content.call_args.args[0]
)
self.assertEqual(content_request.request_body.sequence, 1) self.assertEqual(content_request.request_body.sequence, 1)
def test_edit_message_uses_cardkit_content_for_streaming_card(self): def test_edit_message_uses_cardkit_content_for_streaming_card(self):
@@ -504,7 +564,9 @@ class TestFeishu(unittest.TestCase):
self.assertTrue(success) self.assertTrue(success)
client._api_client.cardkit.v1.card_element.content.assert_called_once() client._api_client.cardkit.v1.card_element.content.assert_called_once()
message_api.patch.assert_not_called() message_api.patch.assert_not_called()
content_request = client._api_client.cardkit.v1.card_element.content.call_args.args[0] content_request = (
client._api_client.cardkit.v1.card_element.content.call_args.args[0]
)
self.assertEqual(content_request.card_id, "card_stream") self.assertEqual(content_request.card_id, "card_stream")
self.assertEqual(content_request.element_id, Feishu.STREAM_CARD_BODY_ELEMENT_ID) self.assertEqual(content_request.element_id, Feishu.STREAM_CARD_BODY_ELEMENT_ID)
self.assertEqual(content_request.request_body.sequence, 1) self.assertEqual(content_request.request_body.sequence, 1)
@@ -545,7 +607,12 @@ class TestFeishu(unittest.TestCase):
{ {
"type": "message", "type": "message",
"text": "", "text": "",
"files": [{"ref": "feishu://file/file_key/report.pdf", "name": "report.pdf"}], "files": [
{
"ref": "feishu://file/file_key/report.pdf",
"name": "report.pdf",
}
],
"message_id": "om_file", "message_id": "om_file",
"chat_id": "oc_chat", "chat_id": "oc_chat",
"sender": { "sender": {
@@ -567,14 +634,18 @@ class TestFeishu(unittest.TestCase):
message_type="image", message_type="image",
content=json.dumps({"image_key": "img_v2_evt"}), content=json.dumps({"image_key": "img_v2_evt"}),
) )
sender = SimpleNamespace(sender_id=SimpleNamespace(open_id="ou_user_evt", user_id=None)) sender = SimpleNamespace(
sender_id=SimpleNamespace(open_id="ou_user_evt", user_id=None)
)
event = SimpleNamespace(sender=sender, message=message) event = SimpleNamespace(sender=sender, message=message)
with patch.object(client, "_forward_to_message_chain") as forward: with patch.object(client, "_forward_to_message_chain") as forward:
client._on_message(SimpleNamespace(event=event)) client._on_message(SimpleNamespace(event=event))
payload = forward.call_args.args[0] payload = forward.call_args.args[0]
self.assertEqual(payload["images"][0]["ref"], "feishu://image/om_img_evt/img_v2_evt") self.assertEqual(
payload["images"][0]["ref"], "feishu://image/om_img_evt/img_v2_evt"
)
def test_on_message_wraps_feishu_audio_ref_with_message_id(self): def test_on_message_wraps_feishu_audio_ref_with_message_id(self):
client = self._build_client() client = self._build_client()
@@ -583,16 +654,23 @@ class TestFeishu(unittest.TestCase):
chat_id="oc_chat_evt", chat_id="oc_chat_evt",
chat_type="p2p", chat_type="p2p",
message_type="audio", message_type="audio",
content=json.dumps({"file_key": "file_audio_evt", "file_name": "voice.opus"}), content=json.dumps(
{"file_key": "file_audio_evt", "file_name": "voice.opus"}
),
)
sender = SimpleNamespace(
sender_id=SimpleNamespace(open_id="ou_user_evt", user_id=None)
) )
sender = SimpleNamespace(sender_id=SimpleNamespace(open_id="ou_user_evt", user_id=None))
event = SimpleNamespace(sender=sender, message=message) event = SimpleNamespace(sender=sender, message=message)
with patch.object(client, "_forward_to_message_chain") as forward: with patch.object(client, "_forward_to_message_chain") as forward:
client._on_message(SimpleNamespace(event=event)) client._on_message(SimpleNamespace(event=event))
payload = forward.call_args.args[0] payload = forward.call_args.args[0]
self.assertEqual(payload["audio_refs"], ["feishu://file/om_audio_evt/file_audio_evt/voice.opus"]) self.assertEqual(
payload["audio_refs"],
["feishu://file/om_audio_evt/file_audio_evt/voice.opus"],
)
def test_feishu_channel_capabilities_enable_images_and_files(self): def test_feishu_channel_capabilities_enable_images_and_files(self):
self.assertTrue( self.assertTrue(
@@ -650,9 +728,12 @@ class TestFeishu(unittest.TestCase):
file_create_response=file_upload_response, file_create_response=file_upload_response,
) )
with tempfile.NamedTemporaryFile(suffix=".txt") as fp, patch.object( with (
tempfile.NamedTemporaryFile(suffix=".txt") as fp,
patch.object(
client, "send_text", return_value={"success": True} client, "send_text", return_value={"success": True}
) as send_text: ) as send_text,
):
fp.write(b"text-bytes") fp.write(b"text-bytes")
fp.flush() fp.flush()
result = client.send_file( result = client.send_file(
@@ -666,7 +747,9 @@ class TestFeishu(unittest.TestCase):
client._api_client.im.v1.file.create.assert_called_once() client._api_client.im.v1.file.create.assert_called_once()
request = message_api.create.call_args.args[0] request = message_api.create.call_args.args[0]
self.assertEqual(request.request_body.msg_type, "file") self.assertEqual(request.request_body.msg_type, "file")
self.assertEqual(json.loads(request.request_body.content)["file_key"], "file_doc") self.assertEqual(
json.loads(request.request_body.content)["file_key"], "file_doc"
)
send_text.assert_called_once() send_text.assert_called_once()
def test_send_voice_uploads_audio_file_and_optionally_sends_caption(self): def test_send_voice_uploads_audio_file_and_optionally_sends_caption(self):
@@ -682,7 +765,9 @@ class TestFeishu(unittest.TestCase):
with tempfile.NamedTemporaryFile(suffix=".opus") as fp: with tempfile.NamedTemporaryFile(suffix=".opus") as fp:
fp.write(b"opus-bytes") fp.write(b"opus-bytes")
fp.flush() fp.flush()
with patch.object(client, "send_text", return_value={"success": True}) as send_text: with patch.object(
client, "send_text", return_value={"success": True}
) as send_text:
result = client.send_voice( result = client.send_voice(
voice_path=fp.name, voice_path=fp.name,
userid="ou_user_8", userid="ou_user_8",
@@ -692,20 +777,30 @@ class TestFeishu(unittest.TestCase):
self.assertTrue(result["success"]) self.assertTrue(result["success"])
request = message_api.create.call_args.args[0] request = message_api.create.call_args.args[0]
self.assertEqual(request.request_body.msg_type, "audio") self.assertEqual(request.request_body.msg_type, "audio")
self.assertEqual(json.loads(request.request_body.content)["file_key"], "file_audio") self.assertEqual(
json.loads(request.request_body.content)["file_key"], "file_audio"
)
send_text.assert_called_once() send_text.assert_called_once()
def test_download_helpers_return_bytes_and_data_url(self): def test_download_helpers_return_bytes_and_data_url(self):
client = self._build_client() client = self._build_client()
client._api_client, _ = self._build_message_api( client._api_client, _ = self._build_message_api(
image_get_response=self._resource_response(b"image-bytes", file_name="poster.png", content_type="image/png"), image_get_response=self._resource_response(
file_get_response=self._resource_response(b"file-bytes", file_name="report.txt", content_type="text/plain"), b"image-bytes", file_name="poster.png", content_type="image/png"
message_resource_response=self._resource_response(b"resource-bytes", file_name="voice.opus", content_type="audio/ogg"), ),
file_get_response=self._resource_response(
b"file-bytes", file_name="report.txt", content_type="text/plain"
),
message_resource_response=self._resource_response(
b"resource-bytes", file_name="voice.opus", content_type="audio/ogg"
),
) )
image_download = client.download_image_bytes("img_v2_test") image_download = client.download_image_bytes("img_v2_test")
file_download = client.download_file_bytes("file_test") file_download = client.download_file_bytes("file_test")
resource_download = client.download_message_resource_bytes("om_test", "file_test", "audio") resource_download = client.download_message_resource_bytes(
"om_test", "file_test", "audio"
)
self.assertEqual(image_download[0], b"image-bytes") self.assertEqual(image_download[0], b"image-bytes")
self.assertEqual(file_download[0], b"file-bytes") self.assertEqual(file_download[0], b"file-bytes")
@@ -722,9 +817,11 @@ class TestFeishu(unittest.TestCase):
"chat_id": "oc_789", "chat_id": "oc_789",
} }
with patch.object(module, "get_configs", return_value={"feishu-main": conf}), patch.object( with (
module, "check_message", return_value=True patch.object(module, "get_configs", return_value={"feishu-main": conf}),
), patch.object(module, "get_instance", return_value=client): patch.object(module, "check_message", return_value=True),
patch.object(module, "get_instance", return_value=client),
):
response = module.send_direct_message( response = module.send_direct_message(
Notification( Notification(
targets={ targets={
@@ -757,14 +854,25 @@ class TestFeishu(unittest.TestCase):
created_loops.append(loop) created_loops.append(loop)
return loop return loop
with patch("app.modules.feishu.feishu.lark_ws_client_module.loop", original_loop), patch( with (
patch(
"app.modules.feishu.feishu.lark_ws_client_module.loop", original_loop
),
patch(
"app.modules.feishu.feishu.lark_ws_client_module._select", "app.modules.feishu.feishu.lark_ws_client_module._select",
new=MagicMock(return_value=None), new=MagicMock(return_value=None),
), patch("app.modules.feishu.feishu.asyncio.new_event_loop", side_effect=_new_loop), patch( ),
patch(
"app.modules.feishu.feishu.asyncio.new_event_loop",
side_effect=_new_loop,
),
patch(
"app.modules.feishu.feishu.lark.ws.Client", return_value=fake_ws_client "app.modules.feishu.feishu.lark.ws.Client", return_value=fake_ws_client
), patch.object( ),
patch.object(
fake_ws_client, "start", side_effect=lambda: None fake_ws_client, "start", side_effect=lambda: None
) as mock_start: ) as mock_start,
):
client._run_ws_client() client._run_ws_client()
self.assertIsNone(client._ws_loop) self.assertIsNone(client._ws_loop)
@@ -784,7 +892,10 @@ class TestFeishu(unittest.TestCase):
future = MagicMock() future = MagicMock()
future.result.return_value = None future.result.return_value = None
with patch("app.modules.feishu.feishu.asyncio.run_coroutine_threadsafe", return_value=future) as runner: with patch(
"app.modules.feishu.feishu.asyncio.run_coroutine_threadsafe",
return_value=future,
) as runner:
client.stop() client.stop()
runner.assert_called_once() runner.assert_called_once()
@@ -795,13 +906,24 @@ class TestFeishu(unittest.TestCase):
client = MagicMock() client = MagicMock()
client.download_image_bytes.return_value = (b"image", "poster.png", "image/png") client.download_image_bytes.return_value = (b"image", "poster.png", "image/png")
client.download_file_bytes.return_value = (b"file", "note.txt", "text/plain") client.download_file_bytes.return_value = (b"file", "note.txt", "text/plain")
client.download_message_resource_bytes.return_value = (b"image", "poster.png", "image/png") client.download_message_resource_bytes.return_value = (
b"image",
"poster.png",
"image/png",
)
with patch.object(module, "get_config", return_value=SimpleNamespace(name="feishu-main")), patch.object( with (
module, "get_instance", return_value=client patch.object(
module, "get_config", return_value=SimpleNamespace(name="feishu-main")
),
patch.object(module, "get_instance", return_value=client),
): ):
data_url = module.download_feishu_image_to_data_url("feishu://image/om_msg/img_v2_xxx", "feishu-main") data_url = module.download_feishu_image_to_data_url(
file_bytes = module.download_feishu_file_bytes("feishu://file/file_xxx/note.txt", "feishu-main") "feishu://image/om_msg/img_v2_xxx", "feishu-main"
)
file_bytes = module.download_feishu_file_bytes(
"feishu://file/file_xxx/note.txt", "feishu-main"
)
audio_bytes = module.download_feishu_file_bytes( audio_bytes = module.download_feishu_file_bytes(
"feishu://file/om_audio/file_audio/voice.opus", "feishu://file/om_audio/file_audio/voice.opus",
"feishu-main", "feishu-main",
@@ -827,11 +949,18 @@ class TestFeishu(unittest.TestCase):
client.add_message_reaction.return_value = "reaction_2" client.add_message_reaction.return_value = "reaction_2"
client.delete_message_reaction.return_value = True client.delete_message_reaction.return_value = True
with patch.object(module, "get_config", return_value=SimpleNamespace(name="feishu-main")), patch.object( with (
module, "get_instance", return_value=client patch.object(
module, "get_config", return_value=SimpleNamespace(name="feishu-main")
),
patch.object(module, "get_instance", return_value=client),
): ):
reaction_id = module.add_feishu_message_reaction("om_x", "GLANCE", "feishu-main") reaction_id = module.add_feishu_message_reaction(
deleted = module.delete_feishu_message_reaction("om_x", "reaction_2", "feishu-main") "om_x", "GLANCE", "feishu-main"
)
deleted = module.delete_feishu_message_reaction(
"om_x", "reaction_2", "feishu-main"
)
self.assertEqual(reaction_id, "reaction_2") self.assertEqual(reaction_id, "reaction_2")
self.assertTrue(deleted) self.assertTrue(deleted)
@@ -842,8 +971,11 @@ class TestFeishu(unittest.TestCase):
client = MagicMock() client = MagicMock()
client.close_streaming_card.return_value = True client.close_streaming_card.return_value = True
with patch.object(module, "get_config", return_value=SimpleNamespace(name="feishu-main")), patch.object( with (
module, "get_instance", return_value=client patch.object(
module, "get_config", return_value=SimpleNamespace(name="feishu-main")
),
patch.object(module, "get_instance", return_value=client),
): ):
success = module.finalize_message( success = module.finalize_message(
MessageResponse( MessageResponse(
@@ -862,18 +994,35 @@ class TestFeishu(unittest.TestCase):
) )
self.assertTrue(success) self.assertTrue(success)
client.close_streaming_card.assert_called_once_with(card_id="card_stream", sequence=3) client.close_streaming_card.assert_called_once_with(
card_id="card_stream", sequence=3
)
def test_module_post_message_prefers_file_and_voice_paths(self): def test_module_post_message_prefers_file_and_voice_paths(self):
module = FeishuModule() module = FeishuModule()
conf = SimpleNamespace(name="feishu-main") conf = SimpleNamespace(name="feishu-main")
client = MagicMock() client = MagicMock()
with patch.object(module, "get_configs", return_value={"feishu-main": conf}), patch.object( with (
module, "check_message", return_value=True patch.object(module, "get_configs", return_value={"feishu-main": conf}),
), patch.object(module, "get_instance", return_value=client): patch.object(module, "check_message", return_value=True),
module.post_message(Notification(file_path="/tmp/demo.txt", text="说明", title="标题", userid="ou_user")) patch.object(module, "get_instance", return_value=client),
module.post_message(Notification(voice_path="/tmp/demo.opus", voice_caption="语音说明", userid="ou_user")) ):
module.post_message(
Notification(
file_path="/tmp/demo.txt",
text="说明",
title="标题",
userid="ou_user",
)
)
module.post_message(
Notification(
voice_path="/tmp/demo.opus",
voice_caption="语音说明",
userid="ou_user",
)
)
client.send_file.assert_called_once() client.send_file.assert_called_once()
client.send_voice.assert_called_once() client.send_voice.assert_called_once()
@@ -883,9 +1032,11 @@ class TestFeishu(unittest.TestCase):
conf = SimpleNamespace(name="feishu-main") conf = SimpleNamespace(name="feishu-main")
client = MagicMock() client = MagicMock()
with patch.object(module, "get_configs", return_value={"feishu-main": conf}), patch.object( with (
module, "check_message", return_value=True patch.object(module, "get_configs", return_value={"feishu-main": conf}),
), patch.object(module, "get_instance", return_value=client): patch.object(module, "check_message", return_value=True),
patch.object(module, "get_instance", return_value=client),
):
module.post_message( module.post_message(
Notification( Notification(
file_path="/tmp/demo.txt", file_path="/tmp/demo.txt",
@@ -917,9 +1068,11 @@ class TestFeishu(unittest.TestCase):
} }
client.send_file.return_value = {"success": True, "message_id": "om_file"} client.send_file.return_value = {"success": True, "message_id": "om_file"}
with patch.object(module, "get_configs", return_value={"feishu-main": conf}), patch.object( with (
module, "check_message", return_value=True patch.object(module, "get_configs", return_value={"feishu-main": conf}),
), patch.object(module, "get_instance", return_value=client): patch.object(module, "check_message", return_value=True),
patch.object(module, "get_instance", return_value=client),
):
response = module.send_direct_message( response = module.send_direct_message(
Notification( Notification(
channel=MessageChannel.Feishu, channel=MessageChannel.Feishu,
@@ -943,9 +1096,11 @@ class TestFeishu(unittest.TestCase):
conf = SimpleNamespace(name="feishu-main") conf = SimpleNamespace(name="feishu-main")
client = MagicMock() client = MagicMock()
with patch.object(module, "get_configs", return_value={"feishu-main": conf}), patch.object( with (
module, "check_message", return_value=True patch.object(module, "get_configs", return_value={"feishu-main": conf}),
), patch.object(module, "get_instance", return_value=client): patch.object(module, "check_message", return_value=True),
patch.object(module, "get_instance", return_value=client),
):
module.post_message( module.post_message(
Notification( Notification(
title="标题", title="标题",
+51 -18
View File
@@ -3,7 +3,7 @@ import sys
import types import types
import unittest import unittest
from pathlib import Path from pathlib import Path
from unittest.mock import call, patch from unittest.mock import patch
def _load_jellyfin_module(): def _load_jellyfin_module():
@@ -58,8 +58,12 @@ def _load_jellyfin_module():
return urljoin(host, path) return urljoin(host, path)
log_module.logger = _Logger() log_module.logger = _Logger()
config_module.settings = types.SimpleNamespace(SUPERUSER="admin", USER_AGENT="MoviePilot") config_module.settings = types.SimpleNamespace(
schemas_module.MediaType = types.SimpleNamespace(MOVIE=types.SimpleNamespace(value="movie")) SUPERUSER="admin", USER_AGENT="MoviePilot"
)
schemas_module.MediaType = types.SimpleNamespace(
MOVIE=types.SimpleNamespace(value="movie")
)
schemas_module.MediaServerItem = object schemas_module.MediaServerItem = object
schemas_module.MediaServerLibrary = object schemas_module.MediaServerLibrary = object
schemas_module.Statistic = object schemas_module.Statistic = object
@@ -90,7 +94,13 @@ def _load_jellyfin_module():
for stub_module in stub_modules.values(): for stub_module in stub_modules.values():
stub_module._jellyfin_test_stub = True stub_module._jellyfin_test_stub = True
jellyfin_path = Path(__file__).resolve().parents[1] / "app" / "modules" / "jellyfin" / "jellyfin.py" jellyfin_path = (
Path(__file__).resolve().parents[1]
/ "app"
/ "modules"
/ "jellyfin"
/ "jellyfin.py"
)
spec = importlib.util.spec_from_file_location(module_name, jellyfin_path) spec = importlib.util.spec_from_file_location(module_name, jellyfin_path)
module = importlib.util.module_from_spec(spec) module = importlib.util.module_from_spec(spec)
assert spec and spec.loader assert spec and spec.loader
@@ -114,9 +124,15 @@ class _FakeResponse:
class JellyfinUserResolutionTest(unittest.TestCase): class JellyfinUserResolutionTest(unittest.TestCase):
def test_loader_does_not_leave_stub_modules_in_sys_modules(self): def test_loader_does_not_leave_stub_modules_in_sys_modules(self):
self.assertNotIn("_test_jellyfin_module", sys.modules) self.assertNotIn("_test_jellyfin_module", sys.modules)
self.assertFalse(getattr(sys.modules.get("app.log"), "_jellyfin_test_stub", False)) self.assertFalse(
self.assertFalse(getattr(sys.modules.get("app.core.config"), "_jellyfin_test_stub", False)) getattr(sys.modules.get("app.log"), "_jellyfin_test_stub", False)
self.assertFalse(getattr(sys.modules.get("app.utils.http"), "_jellyfin_test_stub", False)) )
self.assertFalse(
getattr(sys.modules.get("app.core.config"), "_jellyfin_test_stub", False)
)
self.assertFalse(
getattr(sys.modules.get("app.utils.http"), "_jellyfin_test_stub", False)
)
def _build_client(self) -> Jellyfin: def _build_client(self) -> Jellyfin:
client = Jellyfin.__new__(Jellyfin) client = Jellyfin.__new__(Jellyfin)
@@ -134,9 +150,10 @@ class JellyfinUserResolutionTest(unittest.TestCase):
{"Id": "alice-id", "Name": "alice", "Policy": {"IsAdministrator": False}}, {"Id": "alice-id", "Name": "alice", "Policy": {"IsAdministrator": False}},
] ]
with patch.object(jellyfin_module, "RequestUtils") as request_utils_cls, patch.object( with (
jellyfin_module.logger, "warning" patch.object(jellyfin_module, "RequestUtils") as request_utils_cls,
) as warning_mock: patch.object(jellyfin_module.logger, "warning") as warning_mock,
):
request_utils_cls.return_value.get_res.return_value = _FakeResponse(payload) request_utils_cls.return_value.get_res.return_value = _FakeResponse(payload)
user_id = client.get_user("alice") user_id = client.get_user("alice")
@@ -150,7 +167,10 @@ class JellyfinUserResolutionTest(unittest.TestCase):
{ {
"Id": "visible-admin-id", "Id": "visible-admin-id",
"Name": "visible", "Name": "visible",
"Policy": {"IsAdministrator": True, "EnabledFolders": ["lib-1", "lib-2", "lib-3"]}, "Policy": {
"IsAdministrator": True,
"EnabledFolders": ["lib-1", "lib-2", "lib-3"],
},
}, },
{ {
"Id": "full-admin-id", "Id": "full-admin-id",
@@ -177,14 +197,18 @@ class JellyfinUserResolutionTest(unittest.TestCase):
{ {
"Id": "large-admin-id", "Id": "large-admin-id",
"Name": "large", "Name": "large",
"Policy": {"IsAdministrator": True, "EnabledFolders": ["lib-1", "lib-2", "lib-3"]}, "Policy": {
"IsAdministrator": True,
"EnabledFolders": ["lib-1", "lib-2", "lib-3"],
},
}, },
{"Id": "user-id", "Name": "normal", "Policy": {"IsAdministrator": False}}, {"Id": "user-id", "Name": "normal", "Policy": {"IsAdministrator": False}},
] ]
with patch.object(jellyfin_module, "RequestUtils") as request_utils_cls, patch.object( with (
jellyfin_module.logger, "warning" patch.object(jellyfin_module, "RequestUtils") as request_utils_cls,
) as warning_mock: patch.object(jellyfin_module.logger, "warning") as warning_mock,
):
request_utils_cls.return_value.get_res.return_value = _FakeResponse(payload) request_utils_cls.return_value.get_res.return_value = _FakeResponse(payload)
user_id = client.get_user("admin") user_id = client.get_user("admin")
@@ -193,7 +217,9 @@ class JellyfinUserResolutionTest(unittest.TestCase):
self.assertGreaterEqual(warning_mock.call_count, 2) self.assertGreaterEqual(warning_mock.call_count, 2)
warning_messages = [ warning_messages = [
call.args[0] for call in warning_mock.call_args_list if call.args and isinstance(call.args[0], str) call.args[0]
for call in warning_mock.call_args_list
if call.args and isinstance(call.args[0], str)
] ]
self.assertTrue(any("超级管理员" in message for message in warning_messages)) self.assertTrue(any("超级管理员" in message for message in warning_messages))
self.assertTrue( self.assertTrue(
@@ -205,7 +231,12 @@ class JellyfinUserResolutionTest(unittest.TestCase):
for message in warning_messages for message in warning_messages
) )
) )
self.assertTrue(any(("回退" in message) or ("fallback" in message.lower()) for message in warning_messages)) self.assertTrue(
any(
("回退" in message) or ("fallback" in message.lower())
for message in warning_messages
)
)
def test_get_jellyfin_librarys_returns_empty_when_user_missing(self): def test_get_jellyfin_librarys_returns_empty_when_user_missing(self):
client = self._build_client() client = self._build_client()
@@ -223,7 +254,9 @@ class JellyfinUserResolutionTest(unittest.TestCase):
client.user = "user-id" client.user = "user-id"
with patch.object(jellyfin_module, "RequestUtils") as request_utils_cls: with patch.object(jellyfin_module, "RequestUtils") as request_utils_cls:
request_utils_cls.return_value.get_res.return_value = _FakeResponse({"Items": []}) request_utils_cls.return_value.get_res.return_value = _FakeResponse(
{"Items": []}
)
libraries = client._Jellyfin__get_jellyfin_librarys() libraries = client._Jellyfin__get_jellyfin_librarys()