Files
ucvl-home-vision/baidu.py
T

357 lines
21 KiB
Python
Raw Permalink Normal View History

"""Baidu OAuth and bounded uploads. Secrets stay in a private file, outside DB backups."""
import datetime as dt
import hashlib
import hmac
import http.cookies
import json
import os
import re
import secrets
import threading
import time
import urllib.error
import urllib.parse as url
import urllib.request
BLOCK = 4 * 1024 * 1024
CALLBACK = '/api/baidu/callback'
API = 'https://pan.baidu.com/rest/2.0/xpan/'
OAUTH = 'https://openapi.baidu.com/oauth/2.0/'
class NoRedirect(urllib.request.HTTPRedirectHandler):
def redirect_request(self, *args, **kwargs):
return None
def download_url(value):
p = url.urlsplit(value)
if (p.scheme != 'https' or p.username or p.password or p.port not in (None, 443)
or not any((p.hostname or '').endswith('.'+d) or p.hostname == d
for d in ('baidu.com', 'baidupcs.com'))):
raise ValueError('网盘下载地址未通过校验')
return value
class DownloadRedirect(urllib.request.HTTPRedirectHandler):
def redirect_request(self, req, fp, code, msg, headers, newurl):
download_url(newurl)
return super().redirect_request(req, fp, code, msg, headers, newurl)
class Baidu:
def __init__(self, app):
self.a = app
self.guard = threading.RLock()
self.pending = {}
def read(self):
path = self.a.DATA/'baidu-secrets.json'
if not path.exists():
return dict(config={}, connections={})
return json.loads(path.read_text(encoding='utf8'))
def write(self, value):
path = self.a.DATA/'baidu-secrets.json'
temporary = path.with_suffix('.partial')
with temporary.open('w', encoding='utf8') as stream:
json.dump(value, stream, ensure_ascii=False)
stream.flush(); os.fsync(stream.fileno())
os.chmod(temporary, 0o600)
os.replace(temporary, path)
def key(self, user, purpose='media'):
if purpose == 'backup':
if user['role'] != 'admin':
raise self.a.Problem('数据库云备份仅超级管理员可操作', 403)
return 'shared'
if purpose != 'media':
raise self.a.Problem('网盘用途不正确')
self.a.HOUSEHOLDS.scope(user)
return 'shared'
def public_config(self):
with self.guard:
c = self.read()['config']
return dict(configured=bool(c.get('appKey') and c.get('secretKey')),
appKey=c.get('appKey', ''), hasSecret=bool(c.get('secretKey')),
appFolder=c.get('appFolder', '智家'), redirectUri=c.get('redirectUri', ''),
revision=c.get('revision', 0))
def configure(self, data, user):
a = self.a
if user['role'] != 'admin':
raise a.Problem('开发者应用仅超级管理员可配置', 403)
if set(data) != {'appKey', 'secretKey', 'appFolder', 'redirectUri', 'revision'}:
raise a.Problem('应用配置字段不完整')
with self.guard:
saved = self.read(); old = saved['config']
if type(data['revision']) is not int or data['revision'] != old.get('revision', 0):
raise a.Problem('应用配置已变化,请重新读取', 409)
key = a.clean_text(data['appKey'], 150, True)
secret = a.clean_text(data['secretKey'], 150) or old.get('secretKey', '')
folder = a.clean_text(data['appFolder'], 40, True)
callback = a.clean_text(data['redirectUri'], 400, True)
p = url.urlsplit(callback)
if not secret or not re.fullmatch(r'[A-Za-z0-9]+', key+secret):
raise a.Problem('AppKey 或 SecretKey 格式不正确')
if any(c in folder for c in '/\\:*?"<>|') or folder in ('.', '..'):
raise a.Problem('应用文件夹须与百度后台应用名称一致')
if p.scheme != 'https' or not p.hostname or p.username or p.password or p.query or p.fragment or p.path != CALLBACK:
raise a.Problem('回调地址须为 HTTPS 站点地址加 '+CALLBACK)
changed = any(old.get(k) != v for k, v in dict(appKey=key, secretKey=secret, appFolder=folder, redirectUri=callback).items())
if changed and saved['connections']:
raise a.Problem('已有网盘授权,修改应用前须先解除共享网盘连接', 409)
saved['config'] = dict(appKey=key, secretKey=secret, appFolder=folder,
redirectUri=callback, revision=old.get('revision', 0)+1)
self.write(saved)
if changed: self.pending.clear()
a.audit('配置百度网盘应用', user['username'])
return self.public_config()
def status(self, user, purpose='media'):
key = self.key(user, purpose)
with self.guard:
saved = self.read(); c = saved['config']; connection = saved['connections'].get(key, {})
return dict(configured=bool(c.get('appKey') and c.get('secretKey')),
connected=bool(connection.get('refresh_token')), account=connection.get('name', ''),
reconnect=bool(connection.get('invalid')), purpose=purpose,
redirectUri=c.get('redirectUri', ''), root='/apps/'+c.get('appFolder', '智家'),
connectedAt=connection.get('connectedAt', ''))
def begin(self, user, purpose, session_token, host):
if user['role'] != 'admin':
raise self.a.Problem('共享网盘由超级管理员统一连接', 403)
key = self.key(user, purpose)
with self.guard:
c = self.read()['config']
if not c.get('appKey') or not c.get('secretKey'):
raise self.a.Problem('请先由超级管理员配置百度开发者应用')
if url.urlsplit(c['redirectUri']).netloc != host:
raise self.a.Problem('请通过 '+url.urlsplit(c['redirectUri']).netloc+' 登录系统后连接网盘')
now = time.time()
self.pending = {s: v for s, v in self.pending.items() if v['expires'] > now and v['key'] != key}
if len(self.pending) >= 32:
raise self.a.Problem('授权请求较多,请稍后再试', 429)
state = secrets.token_hex(32)
self.pending[state] = dict(key=key, userId=user['id'], familyId=self.a.HOUSEHOLDS.scope(user),
session=hashlib.sha256(session_token.encode()).hexdigest(),
revision=c['revision'], expires=now+600)
address = OAUTH+'authorize?'+url.urlencode(dict(response_type='code', client_id=c['appKey'],
redirect_uri=c['redirectUri'], scope='basic,netdisk', state=state, display='page'))
return dict(url=address), 'baidu_flow='+state+'; Path=/api/baidu/callback; HttpOnly; SameSite=Lax; Secure; Max-Age=600'
def json_request(self, address, data=None, raw=None, content_type=None):
body = url.urlencode(data).encode() if data is not None else raw
headers = {'User-Agent': 'pan.baidu.com'}
if body is not None: headers['Content-Type'] = content_type or 'application/x-www-form-urlencoded'
req = urllib.request.Request(address, body, headers)
try:
with urllib.request.build_opener(NoRedirect()).open(req, timeout=40) as response:
payload = response.read(2*1024*1024+1)
if len(payload) > 2*1024*1024: raise ValueError()
value = json.loads(payload)
except Exception:
raise self.a.Problem('百度网盘暂时未响应,请稍后重试', 502) from None
if not isinstance(value, dict): raise self.a.Problem('网盘返回内容不完整', 502)
if value.get('error') or value.get('errno', value.get('error_code', 0)):
code = str(value.get('error') or value.get('errno') or value.get('error_code'))
code = code if re.fullmatch('[A-Za-z0-9_-]{1,50}', code) else 'unknown'
if code in ('invalid_grant', 'invalid_token', 'expired_token', '110', '111', '-6', '20016', '20017', '31045'):
error = self.a.Problem('网盘授权已失效,请重新连接('+code+')', 409)
error.reauthorize = True
raise error
raise self.a.Problem('网盘操作未完成('+code+'),请检查应用权限、目录及网盘容量', 502)
return value
def token_response(self, c, **fields):
value = self.json_request(OAUTH+'token?'+url.urlencode(dict(client_id=c['appKey'], client_secret=c['secretKey'], **fields)))
if not value.get('access_token') or not value.get('refresh_token') or not isinstance(value.get('expires_in'), int):
raise self.a.Problem('网盘授权返回不完整,请重新连接', 502)
if 'netdisk' not in value.get('scope', '').split():
# Some responses separate scopes with commas instead of spaces.
if 'netdisk' not in value.get('scope', '').split(','):
raise self.a.Problem('本次未授予网盘访问权限,请重新连接', 403)
return dict(access_token=value['access_token'], refresh_token=value['refresh_token'],
expires=time.time()+max(0, value['expires_in']-120))
def callback(self, query, cookie):
state = query.get('state', '')
try: cookies = http.cookies.SimpleCookie(cookie)
except http.cookies.CookieError: cookies = {}
flow = cookies.get('baidu_flow')
if not re.fullmatch('[a-f0-9]{64}', state) or not flow or not hmac.compare_digest(flow.value, state):
raise self.a.Problem('授权校验失败,请从系统重新发起连接', 403)
with self.guard:
pending = self.pending.pop(state, None)
c = self.read()['config']
if not pending or pending['expires'] < time.time() or pending['revision'] != c.get('revision'):
raise self.a.Problem('授权已过期或已使用,请重新连接', 403)
with self.a.LOCK:
session = self.a.DB.execute('SELECT user_id FROM sessions WHERE hash=? AND expires>?', (pending['session'], time.time())).fetchone()
user = self.a.get_object('users', pending['userId'])
if not session or session[0] != pending['userId'] or not user or user['disabled']:
raise self.a.Problem('原登录已失效,请重新登录后授权', 403)
user = self.a.HOUSEHOLDS.context(user, pending['familyId'])
if user['role'] != 'admin': raise self.a.Problem('管理员权限已变化,请重新授权', 403)
key = self.key(user)
if query.get('error') or not query.get('code'):
raise self.a.Problem('本次未完成网盘授权,可回到系统重新连接')
token = self.token_response(c, grant_type='authorization_code', code=query['code'], redirect_uri=c['redirectUri'])
info = self.json_request(API+'nas?'+url.urlencode(dict(method='uinfo', access_token=token['access_token'])))
uid = str(info.get('uk', ''))
if not uid.isdigit(): raise self.a.Problem('无法确认网盘账号,请重试', 502)
saved = self.read(); old = saved['connections'].get(key, {})
# Re-authorizing a different disk would strand existing attachments.
owner = self.a.setting('baidu_owner:'+key)
if owner and owner != uid:
raise self.a.Problem('请连接原网盘账号;更换账号需先迁移现有资料', 409)
saved['connections'][key] = dict(token, uid=uid, name=str(info.get('baidu_name') or info.get('netdisk_name') or '百度网盘')[:80],
connectedAt=dt.datetime.now(dt.timezone.utc).isoformat(), generation=old.get('generation', 0)+1)
self.write(saved); self.a.set_setting('baidu_owner:'+key, uid)
self.a.audit('连接共享百度网盘', user['username'])
def disconnect(self, user, purpose):
if user['role'] != 'admin': raise self.a.Problem('共享网盘由超级管理员管理', 403)
key = self.key(user, purpose)
with self.guard:
saved = self.read(); saved['connections'].pop(key, None); self.write(saved)
self.pending = {s: v for s, v in self.pending.items() if v['key'] != key}
self.a.audit('解除本机网盘连接', user['username'])
return self.status(user, purpose)
def token(self, key):
with self.guard:
saved = self.read(); connection = saved['connections'].get(key)
if not connection or connection.get('invalid'):
raise self.a.Problem('请连接百度网盘后再操作', 409)
if connection['expires'] <= time.time():
try:
connection.update(self.token_response(saved['config'], grant_type='refresh_token', refresh_token=connection['refresh_token']))
except self.a.Problem as error:
if getattr(error, 'reauthorize', False):
connection['invalid'] = True; self.write(saved)
raise
self.write(saved)
return connection['access_token']
def request(self, key, resource, method, data=None, **query):
return self.json_request(API+resource+'?'+url.urlencode(dict(method=method, access_token=self.token(key), **query)), data=data)
def prepare(self, remote, size, hashes):
root = '/apps/'+self.public_config()['appFolder']
if not remote.startswith(root+'/') or '/..' in remote: raise self.a.Problem('网盘路径不正确')
parent = remote.rsplit('/', 1)[0]
for n in range(3, len(parent.split('/'))+1):
path = '/'.join(parent.split('/')[:n])
# rtype=0 keeps an existing folder; -8 is also reported by some API versions.
try: self.request('shared', 'file', 'create', dict(path=path, isdir=1, rtype=0))
except self.a.Problem as error:
if '(-8)' not in str(error): raise
found = self.find_file(remote, size)
if found: return dict(existing=found)
pre = self.request('shared', 'file', 'precreate', dict(path=remote, size=size, isdir=0, rtype=0, block_list=json.dumps(hashes), autoinit=1))
if pre.get('return_type') == 2:
found = self.find_file(remote, size)
if not found: raise self.a.Problem('云端文件尚未确认,请重试', 502)
return dict(existing=found)
if not pre.get('uploadid'): raise self.a.Problem('网盘预上传返回不完整', 502)
return pre
def find_file(self, remote, size):
values = self.request('shared', 'file', 'list', dir=remote.rsplit('/', 1)[0], limit=1000).get('list', [])
found = next((v for v in values if v.get('path') == remote), None)
if found and found.get('size') != size: raise self.a.Problem('云端同名文件大小不符,未覆盖文件', 409)
return found
def send_block(self, remote, upload_id, index, block, expected):
if hashlib.md5(block).hexdigest() != expected: raise self.a.Problem('文件分片校验不一致,请重新选择原文件')
boundary = secrets.token_hex(16)
body = ('--'+boundary+'\r\nContent-Disposition: form-data; name="file"; filename="chunk"\r\nContent-Type: application/octet-stream\r\n\r\n').encode()+block+('\r\n--'+boundary+'--\r\n').encode()
address = 'https://d.pcs.baidu.com/rest/2.0/pcs/superfile2?'+url.urlencode(dict(method='upload', type='tmpfile', path=remote, uploadid=upload_id, partseq=index, access_token=self.token('shared')))
result = self.json_request(address, raw=body, content_type='multipart/form-data; boundary='+boundary)
if result.get('md5') != expected: raise self.a.Problem('网盘分片校验不一致,请重试', 502)
def complete(self, item):
found = self.find_file(item['remotePath'], item['size'])
if found: return found
result = self.request('shared', 'file', 'create', dict(path=item['remotePath'], size=item['size'], isdir=0, rtype=0,
uploadid=item['uploadId'], block_list=json.dumps(item['hashes'])))
if not result.get('fs_id') or result.get('size') != item['size'] or result.get('path') != item['remotePath']:
raise self.a.Problem('上传结果尚未确认,请重试核对云端文件', 502)
return result
def upload(self, key, local, remote, progress=lambda n, total: None):
"""Only database backups use a local file; personal media uses send_block directly."""
size = local.stat().st_size
hashes = []
with local.open('rb') as stream:
while block := stream.read(BLOCK): hashes.append(hashlib.md5(block).hexdigest())
if not hashes: raise self.a.Problem('空文件不能上传')
pre = self.prepare(remote, size, hashes)
if pre.get('existing'): return pre['existing']
needed = pre.get('block_list', list(range(len(hashes))))
if any(type(i) is not int or not 0 <= i < len(hashes) for i in needed):
raise self.a.Problem('网盘分片索引不正确', 502)
with local.open('rb') as stream:
for i in needed:
stream.seek(i*BLOCK)
self.send_block(remote, pre['uploadid'], i, stream.read(BLOCK), hashes[i])
progress(i+1, len(hashes))
return self.complete(dict(remotePath=remote, size=size, uploadId=pre['uploadid'], hashes=hashes))
def open_file(self, key, item, byte_range=''):
result = self.request(key, 'multimedia', 'filemetas', fsids=json.dumps([int(item['fsId'])]), dlink=1)
info = next((v for v in result.get('list', []) if str(v.get('fs_id')) == str(item['fsId'])), None)
if not info or info.get('path') != item['remotePath']:
raise self.a.Problem('网盘文件已移动或不存在', 404)
address = download_url(info.get('dlink', ''))
address += ('&' if '?' in address else '?')+url.urlencode(dict(access_token=self.token(key)))
headers = {'User-Agent': 'pan.baidu.com'}
if byte_range:
if not re.fullmatch(r'bytes=(?:\d+-\d*|-\d+)', byte_range): raise self.a.Problem('播放范围不正确', 416)
headers['Range'] = byte_range
try:
return urllib.request.build_opener(DownloadRedirect()).open(urllib.request.Request(address, headers=headers), timeout=30)
except Exception:
raise self.a.Problem('网盘文件暂时无法读取,请稍后重试', 502) from None
def verify_backup(self, remote, fs_id, size, digest):
"""Read back into a small memory buffer before allowing retention deletion."""
sha = hashlib.sha256(); received = 0
with self.open_file('shared', dict(remotePath=remote, fsId=str(fs_id))) as response:
while chunk := response.read(65536):
received += len(chunk)
if received > size: raise self.a.Problem('云端备份大小不符,未清理旧备份', 502)
sha.update(chunk)
if received != size or sha.hexdigest() != digest:
raise self.a.Problem('云端备份校验不一致,未清理旧备份', 502)
def list_directory(self, directory):
items = []; seen = set()
for start in range(0, 20000, 1000):
result = self.request('shared', 'file', 'list', dir=directory, start=start, limit=1000, order='name')
page = result.get('list')
if not isinstance(page, list): raise self.a.Problem('网盘目录读取不完整,未继续清理', 502)
for item in page:
path = item.get('path', '')
if path in seen: raise self.a.Problem('网盘目录分页重复,未继续清理', 502)
seen.add(path); items.append(item)
if len(page) < 1000: return items
raise self.a.Problem('网盘目录过大,未继续清理', 502)
def delete_backup_file(self, remote):
# Restrict this capability to one generated database file, never a directory or media.
root = '/apps/'+self.public_config()['appFolder']+'/数据库备份/'
if not remote.startswith(root): raise self.a.Problem('备份清理路径不正确')
parts = remote[len(root):].split('/')
if len(parts) not in (2, 3) or not re.fullmatch('[a-f0-9]{16}', parts[0]):
raise self.a.Problem('备份清理路径不正确')
name = parts[-1]
if not re.fullmatch('vision-'+parts[0]+r'-\d{8}T\d{12}Z\.sqlite', name) or len(parts) == 3 and parts[1] != name:
raise self.a.Problem('备份清理文件名不正确')
result = self.request('shared', 'file', 'filemanager', {'async':0, 'filelist':json.dumps([remote])}, opera='delete')
info = result.get('info')
if result.get('taskid') or not isinstance(info, list) or len(info) != 1 or info[0].get('path') != remote or info[0].get('errno') != 0:
raise self.a.Problem('网盘未确认删除结果,请稍后重试', 502)