Files
video-gen/video-gen-api/app/utils/douyinRequest.py
T

376 lines
14 KiB
Python

import json
import asyncio
from typing import Any, Dict, Optional, Tuple
import httpx
from sqlalchemy import select, update
from sqlalchemy.ext.asyncio import AsyncSession
from datetime import datetime, timezone, timedelta
from app.models.user_oauth import UserOAuth
from app.models.user_oauth_app import UserOAuthApp
from app.models.base import async_session
from app.utils.redis import get_redis
from app.utils.logger import get_logger
logger = get_logger("douyin_request", "douyin_request")
class DouyinRequest:
def __init__(self, platform: str = "douyin"):
self._client: Optional[httpx.AsyncClient] = None
self._platform = platform
@property
def client(self) -> httpx.AsyncClient:
if not self._client:
self._client = httpx.AsyncClient(
timeout=httpx.Timeout(30.0),
follow_redirects=True,
)
return self._client
@property
def _redis_key(self) -> str:
return f"{self._platform}:tokens"
async def close(self):
if self._client:
await self._client.aclose()
self._client = None
async def _get_redis_token(self, oauth_id: str) -> Optional[str]:
redis = get_redis()
if not redis:
raise ValueError("Redis连接未配置")
try:
cache_str = await redis.hget(self._redis_key, oauth_id)
if cache_str:
cache = json.loads(cache_str)
token = cache.get("token")
expired_at_str = cache.get("expired_at")
if token and expired_at_str:
expired_at = datetime.fromisoformat(expired_at_str)
if expired_at > datetime.now(timezone.utc):
return token
except Exception as e:
raise ValueError(f"获取Redis缓存失败: {e}")
return None
async def _set_redis_token(self, oauth_id: str, token: str, expired_at: datetime):
redis = get_redis()
if not redis:
raise ValueError("Redis连接未配置")
try:
cache = {
"token": token,
"expired_at": expired_at.isoformat(),
}
await redis.hset(self._redis_key, oauth_id, json.dumps(cache))
except Exception as e:
raise ValueError(f"设置Redis缓存失败: {e}")
async def _delete_redis_token(self, oauth_id: str):
redis = get_redis()
if not redis:
raise ValueError("Redis连接未配置")
try:
await redis.hdel(self._redis_key, oauth_id)
except Exception as e:
raise ValueError(f"删除Redis缓存失败: {e}")
async def get_access_token(self, oauth_id: str, force_refresh: bool = False) -> str:
if not force_refresh:
token = await self._get_redis_token(oauth_id)
if token:
return token
async with async_session() as db:
oauth_data = await db.execute(
select(UserOAuth).where(
UserOAuth.id == oauth_id,
UserOAuth.deleted_at.is_(None),
).limit(1)
)
oauth_data = oauth_data.scalar_one_or_none()
if not oauth_data:
raise ValueError("无效的oauth_id")
if force_refresh:
where_cond = UserOAuth.deleted_at.is_(None)
if oauth_data.appid:
where_cond = where_cond & (UserOAuth.appid == oauth_data.appid)
if oauth_data.account_username:
where_cond = where_cond & (UserOAuth.account_username == oauth_data.account_username)
if oauth_data.account_userid:
where_cond = where_cond & (UserOAuth.account_userid == oauth_data.account_userid)
await db.execute(
update(UserOAuth).where(where_cond).values(
access_token=None,
access_token_expired=None,
)
)
await db.commit()
related_oauth_ids = await db.execute(
select(UserOAuth.id).where(where_cond)
)
related_oauth_ids = [row[0] for row in related_oauth_ids.all()]
for related_id in related_oauth_ids:
await self._delete_redis_token(related_id)
new_token, new_expired_at = await self.refresh_access_token(db, oauth_id, oauth_data.appid, oauth_data.refresh_token)
return new_token
if oauth_data.access_token_expired and oauth_data.access_token_expired > datetime.now(timezone.utc):
token = oauth_data.access_token
expired_at = oauth_data.access_token_expired
await self._set_redis_token(oauth_id, token, expired_at)
return token
where_cond = UserOAuth.deleted_at.is_(None)
if oauth_data.appid:
where_cond = where_cond & (UserOAuth.appid == oauth_data.appid)
if oauth_data.account_username:
where_cond = where_cond & (UserOAuth.account_username == oauth_data.account_username)
if oauth_data.account_userid:
where_cond = where_cond & (UserOAuth.account_userid == oauth_data.account_userid)
related_oauths = await db.execute(
select(UserOAuth.access_token, UserOAuth.access_token_expired).where(
where_cond,
UserOAuth.access_token_expired.is_not(None),
UserOAuth.access_token_expired > datetime.now(timezone.utc),
).limit(1)
)
related_oauth = related_oauths.first()
if related_oauth:
token, expired_at = related_oauth
await self._set_redis_token(oauth_id, token, expired_at)
return token
if oauth_data.refresh_token_expired and oauth_data.refresh_token_expired < datetime.now(timezone.utc):
raise ValueError("授权已过期,请重新授权")
new_token, new_expired_at = await self.refresh_access_token(db, oauth_id, oauth_data.appid, oauth_data.refresh_token)
return new_token
async def refresh_access_token(self, db: AsyncSession, oauth_id: str, appid: str, refresh_token: str) -> Tuple[str, datetime]:
oauth_info = await db.execute(
select(UserOAuth.account_username, UserOAuth.account_userid).where(
UserOAuth.id == oauth_id,
UserOAuth.deleted_at.is_(None),
).limit(1)
)
oauth_info = oauth_info.first()
if not oauth_info:
raise ValueError("无效的oauth_id")
account_username, account_userid = oauth_info
result = await db.execute(
select(UserOAuthApp.secret).where(
UserOAuthApp.app_id == appid,
UserOAuthApp.status == 1,
UserOAuthApp.deleted_at.is_(None),
).limit(1)
)
app_secret = result.scalar_one_or_none()
if not app_secret:
raise ValueError("应用已被删除或禁用")
response = await self.client.request(
'POST',
'https://api.oceanengine.com/open_api/oauth2/refresh_token/',
data={
'app_id': appid,
'secret': app_secret,
'refresh_token': refresh_token,
},
)
response.raise_for_status()
data = response.json()
if 'code' not in data:
raise ValueError("刷新access_token失败,接口未返回code")
code = data.get('code', 0)
if code != 0:
raise ValueError(f"刷新access_token失败,接口返回:{data}")
data = data.get('data', {})
new_access_token = data.get('access_token', '')
new_refresh_token = data.get('refresh_token', '')
expires_in = data.get('expires_in', 0)
refresh_token_expires_in = data.get('refresh_token_expires_in', 0)
new_expired_at = datetime.now(timezone.utc) + timedelta(seconds=expires_in)
where_cond = UserOAuth.deleted_at.is_(None)
if appid:
where_cond = where_cond & (UserOAuth.appid == appid)
if account_username:
where_cond = where_cond & (UserOAuth.account_username == account_username)
if account_userid:
where_cond = where_cond & (UserOAuth.account_userid == account_userid)
await db.execute(
update(UserOAuth).where(
where_cond,
).values(
access_token=new_access_token,
refresh_token=new_refresh_token,
access_token_expired=new_expired_at,
refresh_token_expired=datetime.now(timezone.utc) + timedelta(seconds=refresh_token_expires_in),
)
)
await db.commit()
related_oauth_ids = await db.execute(
select(UserOAuth.id).where(where_cond)
)
related_oauth_ids = [row[0] for row in related_oauth_ids.all()]
for related_id in related_oauth_ids:
await self._set_redis_token(related_id, new_access_token, new_expired_at)
return new_access_token, new_expired_at
# 有token请求
async def request_with_token_with_context(
self,
oauth_id: str,
url: str,
method: str = 'GET',
options: any = None,
request_count: int = 1,
) -> Any:
options = options or {}
token = await self.get_access_token(oauth_id)
for i in range(1, request_count+1):
# 有一些错误是触发频次管理的,需要重试
try:
headers = options.get('headers', {}).copy()
headers['Access-Token'] = token
has_files = 'files' in options
if has_files:
# 移除可能错误设置的 Content-Type,让库自动生成 multipart 头
headers.pop('Content-Type', None)
else:
headers.setdefault('Content-Type', 'application/json')
options['headers'] = headers
response = await self.client.request(method, url, **options)
response.raise_for_status()
try:
data = response.json()
except json.JSONDecodeError:
data = {'code': 0, 'data': response.text, 'msg':'JSON解析失败'}
if 'code' not in data:
await asyncio.sleep(i * 5)
continue
code = data.get('code', 0)
if code in [40102, 40104]:
await asyncio.sleep(i * 5)
token = await self.get_access_token(oauth_id, force_refresh=True)
continue
if code in [40100, 40110]:
wait_time = min(2 * (2 ** (i - 1)), 10)
await asyncio.sleep(wait_time)
continue
if code == 50000:
await asyncio.sleep(i * 10)
continue
return data
except httpx.HTTPStatusError as e:
await asyncio.sleep(i * 10)
continue
except httpx.RequestError as e:
await asyncio.sleep(i * 10)
continue
options_log = {}
if options:
for key, value in options.items():
if key == 'files':
options_log[key] = {k: (v[0], 'bytes_content', v[2]) for k, v in value.items()}
else:
options_log[key] = value
res = json.dumps(data, ensure_ascii=False) if 'data' in locals() else ''
logger.error(
f'DouYin API request failed after {request_count} retries. '
f'url:{url};method:{method};oauth_id:{oauth_id};options:{json.dumps(options_log, ensure_ascii=False)};response:{res}'
)
if 'data' in locals() and data.get('code', 0) != 0:
raise ValueError(f'接口返回错误[code:{data.get("code", "接口编码")}]{data.get("message", "接口返回错误")}')
else:
raise ValueError('网络错误,稍后重试。')
# 无token请求
async def request_with_context(
self,
url: str,
method: str = 'GET',
options: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
options = options or {}
for i in range(1, 2):
try:
headers = options.get('headers', {}).copy()
headers.setdefault('Content-Type', 'application/json')
options['headers'] = headers
response = await self.client.request(method, url, **options)
response.raise_for_status()
try:
data = response.json()
except json.JSONDecodeError:
data = {'code': 0, 'data': response.text}
if 'code' not in data:
await asyncio.sleep(i * 5)
continue
code = data.get('code', 0)
if code >= 50000:
await asyncio.sleep(i * 10)
continue
return data
except httpx.HTTPStatusError as e:
await asyncio.sleep(i * 10)
continue
except httpx.RequestError as e:
await asyncio.sleep(i * 10)
continue
res = json.dumps(data) if 'data' in locals() else ''
raise RuntimeError(
f'DouYin API request failed after 5 retries. '
f'url:{url};options:{json.dumps(options)};response:{res}'
)