359 lines
14 KiB
Python
359 lines
14 KiB
Python
import random
|
|
from datetime import datetime, timedelta
|
|
|
|
import httpx
|
|
from sqlalchemy import select, func, update
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
# 引入项目统一北京时间时区
|
|
from app.utils.datetime_util import BEIJING_TZ
|
|
from app.models.user_oauth import UserOAuth
|
|
from app.models.user_oauth_app import UserOAuthApp
|
|
from app.utils.id_gen import generate_id
|
|
from app.tasks.token_refresh_task import _update_redis_token, _delete_redis_token
|
|
|
|
# 随机获取一个可用的应用配置
|
|
async def get_available_app(open_type: int, db: AsyncSession) -> dict:
|
|
if not isinstance(open_type, int):
|
|
raise ValueError("open_type必须为整数类型")
|
|
|
|
result = await db.execute(
|
|
select(UserOAuthApp).where(
|
|
UserOAuthApp.open_type == open_type,
|
|
UserOAuthApp.status == 1,
|
|
UserOAuthApp.deleted_at.is_(None),
|
|
)
|
|
)
|
|
apps = result.scalars().all()
|
|
|
|
if not apps:
|
|
raise ValueError("没有找到可用的应用配置")
|
|
|
|
available_apps = []
|
|
for app in apps:
|
|
app_id = app.app_id
|
|
count_result = await db.execute(
|
|
select(func.count(UserOAuth.id)).where(
|
|
UserOAuth.appid == app_id,
|
|
UserOAuth.deleted_at.is_(None),
|
|
)
|
|
)
|
|
count = count_result.scalar() or 0
|
|
max_users = app.max_count
|
|
if count < max_users:
|
|
available_apps.append({
|
|
"app_id": app.app_id,
|
|
"secret": app.secret,
|
|
"open_type": app.open_type,
|
|
"auth_url": app.auth_url,
|
|
"company": app.company,
|
|
})
|
|
|
|
if not available_apps:
|
|
raise ValueError("所有应用授权已超过最大数量")
|
|
|
|
return random.choice(available_apps)
|
|
|
|
|
|
async def build_oauth_url(open_type: int, user_id: str, db: AsyncSession) -> str:
|
|
if open_type == 1:
|
|
return await _build_jl_oauth_url(open_type, user_id, db)
|
|
elif open_type in [2, 3]:
|
|
return await _build_jl_oauth_url(2, user_id, db)
|
|
elif open_type == 5:
|
|
return "无配置"
|
|
elif open_type == 9:
|
|
return "无配置"
|
|
elif open_type == 10:
|
|
return "无配置"
|
|
else:
|
|
# 修复BUG:变量名错误 app_type -> open_type
|
|
raise ValueError(f"不支持的应用类型: {open_type}")
|
|
|
|
|
|
# 千川授权链接构建
|
|
async def _build_jl_oauth_url(open_type: int, user_id: str, db: AsyncSession) -> str:
|
|
app = await get_available_app(open_type, db)
|
|
app_id = app.get("app_id")
|
|
auth_url = app.get("auth_url")
|
|
if not auth_url:
|
|
raise ValueError("应用授权链接不能为空")
|
|
params = {
|
|
"app_id": app_id,
|
|
"state": f"{user_id}:{app_id}",
|
|
}
|
|
query_string = "&".join(f"{k}={v}" for k, v in params.items())
|
|
return f"{auth_url}&{query_string}"
|
|
|
|
|
|
async def get_token(code: str, user_id: str, app_id: str, db: AsyncSession) -> dict:
|
|
if not all([code, user_id, app_id]):
|
|
raise ValueError("code、user_id、app_id不能为空")
|
|
|
|
result = await db.execute(
|
|
select(UserOAuthApp).where(
|
|
UserOAuthApp.app_id == app_id,
|
|
UserOAuthApp.deleted_at.is_(None),
|
|
UserOAuthApp.status == 1
|
|
)
|
|
)
|
|
app = result.scalar_one_or_none()
|
|
if not app:
|
|
raise ValueError("应用配置不存在或已禁用")
|
|
|
|
open_type = app.open_type
|
|
secret = app.secret
|
|
|
|
if open_type in (1, 2):
|
|
return await get_juliang_token(app_id, secret, code, open_type, user_id, db)
|
|
elif open_type == 5:
|
|
return await get_kuaishou_token(app_id, secret, code, open_type, user_id, db)
|
|
elif open_type in (9, 10):
|
|
return await get_tencent_token(app_id, secret, code, open_type, user_id, db)
|
|
else:
|
|
raise ValueError(f"不支持的应用类型: {open_type}")
|
|
|
|
|
|
async def get_juliang_token(app_id: str, secret: str, code: str, open_type: int, user_id: str, db: AsyncSession) -> dict:
|
|
timeout = httpx.Timeout(30.0)
|
|
try:
|
|
async with httpx.AsyncClient(timeout=timeout) as client:
|
|
# 1. 获取token
|
|
url = "https://api.oceanengine.com/open_api/oauth2/access_token/"
|
|
resp = await client.post(
|
|
url,
|
|
json={
|
|
"app_id": app_id,
|
|
"secret": secret,
|
|
"auth_code": code,
|
|
},
|
|
)
|
|
resp.raise_for_status()
|
|
data = resp.json()
|
|
if data.get("code") != 0:
|
|
msg = data.get("message", "获取token失败")
|
|
raise ValueError(f"获取token失败:{msg}")
|
|
|
|
resp_data = data.get("data", {})
|
|
access_token = resp_data.get("access_token", "")
|
|
refresh_token = resp_data.get("refresh_token", "")
|
|
expires_sec = resp_data.get("expires_in", 7200)
|
|
refresh_expires_sec = resp_data.get("refresh_token_expires_in", 30 * 24 * 3600)
|
|
|
|
# 【修复:统一使用北京时间计算过期时间】
|
|
now_beijing = datetime.now(tz=BEIJING_TZ)
|
|
access_expired = now_beijing + timedelta(seconds=expires_sec)
|
|
refresh_expired = now_beijing + timedelta(seconds=refresh_expires_sec)
|
|
|
|
# 2. 获取授权广告账户列表
|
|
resp = await client.get(
|
|
"https://api.oceanengine.com/open_api/oauth2/advertiser/get/",
|
|
headers={"Access-Token": access_token},
|
|
)
|
|
resp.raise_for_status()
|
|
data = resp.json()
|
|
if data.get("code") != 0:
|
|
raise ValueError(f"获取已授权账户失败:{data.get('message')}")
|
|
account_list = data.get("data", {}).get("list", [])
|
|
|
|
# 3. 获取登录用户信息
|
|
resp = await client.get(
|
|
"https://api.oceanengine.com/open_api/2/user/info/",
|
|
headers={"Access-Token": access_token},
|
|
)
|
|
resp.raise_for_status()
|
|
data = resp.json()
|
|
if data.get("code") != 0:
|
|
raise ValueError(f"获取登录信息失败:{data.get('message')}")
|
|
|
|
user_info = data.get("data", {})
|
|
account_username = user_info.get("email", "")
|
|
account_userid = str(user_info.get("id", ""))
|
|
material_auth_status = user_info.get("material_auth_status", False)
|
|
|
|
# 4. 查询当前用户+应用+登录账号下的授权记录
|
|
oauth_query = select(UserOAuth).where(
|
|
UserOAuth.account_username == account_username,
|
|
UserOAuth.account_userid == account_userid,
|
|
UserOAuth.appid == app_id,
|
|
UserOAuth.user_id == user_id,
|
|
UserOAuth.deleted_at.is_(None)
|
|
)
|
|
oauth_exist = await db.execute(oauth_query.limit(1))
|
|
oauth_exist = oauth_exist.scalar_one_or_none()
|
|
|
|
if not oauth_exist:
|
|
# 新增授权
|
|
new_oauth_ids = []
|
|
for account in account_list:
|
|
oauth_id = generate_id()
|
|
new_oauth_ids.append(oauth_id)
|
|
db.add(UserOAuth(
|
|
id=oauth_id,
|
|
account_id=str(account.get("account_id", "")),
|
|
account_name=account.get("account_name", ""),
|
|
account_role=account.get("account_role", ""),
|
|
account_username=account_username,
|
|
account_userid=account_userid,
|
|
user_id=user_id,
|
|
appid=app_id,
|
|
open_type=open_type,
|
|
port_type=1,
|
|
access_token=access_token,
|
|
access_token_expired=access_expired,
|
|
refresh_token=refresh_token,
|
|
refresh_token_expired=refresh_expired,
|
|
material_auth_status=material_auth_status,
|
|
))
|
|
await db.commit()
|
|
# 批量写入Redis
|
|
for oauth_id in new_oauth_ids:
|
|
await _update_redis_token(oauth_id, access_token, access_expired)
|
|
else:
|
|
# 查询当前所有有效授权账户
|
|
exist_query = select(UserOAuth).where(
|
|
UserOAuth.account_username == account_username,
|
|
UserOAuth.account_userid == account_userid,
|
|
UserOAuth.appid == app_id,
|
|
UserOAuth.user_id == user_id,
|
|
UserOAuth.deleted_at.is_(None),
|
|
)
|
|
exist_res = await db.execute(exist_query)
|
|
exist_map = {item.account_id: item for item in exist_res.scalars().all()}
|
|
|
|
new_account_ids = {str(acc.get("account_id")) for acc in account_list}
|
|
old_account_ids = set(exist_map.keys())
|
|
update_redis_ids = []
|
|
|
|
# 下线的账户:软删除 + 清理Redis脏缓存
|
|
for del_account_id in old_account_ids - new_account_ids:
|
|
del_oauth = exist_map[del_account_id]
|
|
del_oauth.deleted_at = now_beijing
|
|
await _delete_redis_token(del_oauth.id)
|
|
|
|
# 更新/新增当前授权账户
|
|
for account in account_list:
|
|
aid = str(account.get("account_id", ""))
|
|
if aid in exist_map:
|
|
item = exist_map[aid]
|
|
item.account_name = account.get("account_name", "")
|
|
item.account_role = account.get("account_role", "")
|
|
item.access_token = access_token
|
|
item.access_token_expired = access_expired
|
|
item.refresh_token = refresh_token
|
|
item.refresh_token_expired = refresh_expired
|
|
item.material_auth_status = material_auth_status
|
|
update_redis_ids.append(item.id)
|
|
else:
|
|
oauth_id = generate_id()
|
|
update_redis_ids.append(oauth_id)
|
|
db.add(UserOAuth(
|
|
id=oauth_id,
|
|
account_id=aid,
|
|
account_name=account.get("account_name", ""),
|
|
account_role=account.get("account_role", ""),
|
|
account_username=account_username,
|
|
account_userid=account_userid,
|
|
user_id=user_id,
|
|
appid=app_id,
|
|
open_type=open_type,
|
|
port_type=1,
|
|
access_token=access_token,
|
|
access_token_expired=access_expired,
|
|
refresh_token=refresh_token,
|
|
refresh_token_expired=refresh_expired,
|
|
material_auth_status=material_auth_status,
|
|
))
|
|
await db.commit()
|
|
# 更新有效账号缓存
|
|
for oid in update_redis_ids:
|
|
await _update_redis_token(oid, access_token, access_expired)
|
|
|
|
# 5. 同登录账号、同应用、其他用户下的授权批量同步最新token
|
|
related_query = select(UserOAuth).where(
|
|
UserOAuth.account_username == account_username,
|
|
UserOAuth.account_userid == account_userid,
|
|
UserOAuth.appid == app_id,
|
|
UserOAuth.user_id != user_id,
|
|
UserOAuth.deleted_at.is_(None),
|
|
)
|
|
related_res = await db.execute(related_query)
|
|
related_list = related_res.scalars().all()
|
|
if related_list:
|
|
await db.execute(
|
|
update(UserOAuth).where(
|
|
UserOAuth.account_username == account_username,
|
|
UserOAuth.account_userid == account_userid,
|
|
UserOAuth.appid == app_id,
|
|
UserOAuth.user_id != user_id,
|
|
UserOAuth.deleted_at.is_(None),
|
|
).values(
|
|
access_token=access_token,
|
|
access_token_expired=access_expired,
|
|
refresh_token=refresh_token,
|
|
refresh_token_expired=refresh_expired,
|
|
material_auth_status=material_auth_status,
|
|
)
|
|
)
|
|
await db.commit()
|
|
for item in related_list:
|
|
await _update_redis_token(item.id, access_token, access_expired)
|
|
|
|
return {"message": "授权成功"}
|
|
|
|
except httpx.HTTPError as e:
|
|
raise ValueError(f"第三方接口请求异常:{str(e)}")
|
|
except Exception as e:
|
|
raise ValueError(f"授权处理异常:{str(e)}")
|
|
|
|
|
|
async def get_kuaishou_token(app_id: str, secret: str, code: str, oauth_type: int, user_id: str, db: AsyncSession) -> dict:
|
|
return {"message": "快手渠道暂未实现授权逻辑"}
|
|
|
|
|
|
async def get_tencent_token(app_id: str, secret: str, code: str, oauth_type: int, user_id: str, db: AsyncSession) -> str:
|
|
return "腾讯营销渠道暂未实现授权逻辑"
|
|
|
|
|
|
async def get_oauth_list(
|
|
user_id: str,
|
|
db: AsyncSession,
|
|
account_userid: str | None = None,
|
|
open_type: int | None = None,
|
|
account_id: str | None = None,
|
|
page: int = 1,
|
|
page_size: int = 10,
|
|
) -> dict:
|
|
# 分页参数容错
|
|
page = max(page, 1)
|
|
page_size = max(min(page_size, 100), 1)
|
|
|
|
base_where = [
|
|
UserOAuth.user_id == user_id,
|
|
UserOAuth.deleted_at.is_(None),
|
|
]
|
|
if account_userid:
|
|
base_where.append(UserOAuth.account_userid == account_userid)
|
|
if open_type:
|
|
base_where.append(UserOAuth.open_type == open_type)
|
|
if account_id:
|
|
base_where.append(UserOAuth.account_id == account_id)
|
|
|
|
# 统计总条数(优化:使用count,避免全量查询)
|
|
count_stmt = select(func.count(UserOAuth.id)).where(*base_where)
|
|
total = await db.scalar(count_stmt) or 0
|
|
|
|
# 分页查询数据
|
|
data_stmt = select(UserOAuth).where(*base_where)\
|
|
.order_by(UserOAuth.created_at.desc())\
|
|
.offset((page - 1) * page_size)\
|
|
.limit(page_size)
|
|
result = await db.execute(data_stmt)
|
|
data_list = result.scalars().all()
|
|
|
|
return {
|
|
"data": data_list,
|
|
"total": total,
|
|
"page": page,
|
|
"page_size": page_size,
|
|
} |