ForcePilot/backend/package/yuxi/channels/adapters/msteams/oauth_flow.py
Kris 939f1ba82a refactor(msteams): 整理代码结构并新增多项功能
本次提交对Microsoft Teams适配器代码进行了多维度优化与新增:
1.  调整多处导入顺序,优化代码可读性
2.  新增media_tools工具模块,提供媒体相关辅助函数
3.  新增thread_history模块,实现对话历史拉取与缓存功能
4.  新增connection_modes模块,支持webhook/websocket/polling三种连接模式
5.  扩展security.py与tool_policy.py,新增通配符配置校验与三级策略解析
6.  新增feedback会话记录功能
7.  为sent_message_cache添加自动清理任务
8.  优化normalizer模块,新增引用、编辑消息解析与线程上下文注入
9.  重构file_upload的SSRF防护逻辑,复用公共校验工具
10. 修复多处导入顺序与代码排版问题
11. 为消息发送添加断路器保护与异步去重锁
2026-05-13 16:12:31 +08:00

285 lines
9.8 KiB
Python

"""Microsoft Teams OAuth PKCE 委托授权流程。
实现 OAuth 2.0 Authorization Code + PKCE S256 完整流程,
包含本地回调服务器、浏览器授权引导、WSL2/SSH 检测与手动模式。
"""
from __future__ import annotations
import asyncio
import base64
import hashlib
import os
import secrets
import sys
import webbrowser
from http.server import BaseHTTPRequestHandler, HTTPServer
from typing import Any
from urllib.parse import parse_qs, urlencode, urlparse
import aiohttp
from yuxi.utils.logging_config import logger
AUTHORIZE_URL = "https://login.microsoftonline.com/{tenant}/oauth2/v2.0/authorize"
TOKEN_URL = "https://login.microsoftonline.com/{tenant}/oauth2/v2.0/token"
DEFAULT_CALLBACK_PORT = 5353
DEFAULT_CALLBACK_PATH = "/oauth/msteams/callback"
PKCE_VERIFIER_LENGTH = 64
STATE_LENGTH = 32
def _generate_pkce_code_verifier() -> str:
return base64.urlsafe_b64encode(secrets.token_bytes(PKCE_VERIFIER_LENGTH)).rstrip(b"=").decode()
def _generate_pkce_code_challenge(verifier: str) -> str:
digest = hashlib.sha256(verifier.encode()).digest()
return base64.urlsafe_b64encode(digest).rstrip(b"=").decode()
def _generate_state() -> str:
return secrets.token_hex(STATE_LENGTH)
def detect_wsl2() -> bool:
if sys.platform != "win32":
return False
try:
result = os.popen("wsl.exe --status 2>&1").read()
return "Default Distribution" in result or "Default Version" in result
except Exception:
pass
try:
with open("/proc/version") as f:
content = f.read()
return "microsoft" in content.lower() or "wsl" in content.lower()
except Exception:
pass
return False
def detect_ssh_session() -> bool:
return bool(os.environ.get("SSH_TTY") or os.environ.get("SSH_CONNECTION") or os.environ.get("SSH_CLIENT"))
def needs_manual_oauth() -> bool:
wsl = detect_wsl2()
ssh = detect_ssh_session()
if wsl:
logger.info("MSTeams OAuth: WSL2 detected, using manual auth mode")
if ssh:
logger.info("MSTeams OAuth: SSH session detected, using manual auth mode")
return wsl or ssh
class _CallbackHandler(BaseHTTPRequestHandler):
callback_result: dict[str, str] = {}
def do_GET(self):
parsed = urlparse(self.path)
if parsed.path == DEFAULT_CALLBACK_PATH:
qs = parse_qs(parsed.query)
code = qs.get("code", [""])[0]
state = qs.get("state", [""])[0]
error = qs.get("error", [""])[0]
error_desc = qs.get("error_description", [""])[0]
self.send_response(200)
self.send_header("Content-Type", "text/html; charset=utf-8")
self.end_headers()
if error:
_CallbackHandler.callback_result = {"error": error, "error_description": error_desc}
self.wfile.write(
b"<html><body><h2>Authorization Failed</h2><p>"
+ error_desc.encode()
+ b"</p><p>You can close this window.</p></body></html>"
)
else:
_CallbackHandler.callback_result = {"code": code, "state": state}
self.wfile.write(
b"<html><body><h2>Authorization Successful!</h2>"
b"<p>You can close this window and return to the terminal.</p></body></html>"
)
else:
self.send_response(404)
self.end_headers()
self.wfile.write(b"Not Found")
def log_message(self, format, *args):
pass
class OAuthPKCEFlow:
"""OAuth 2.0 Authorization Code + PKCE S256 完整流程。
支持自动本地回调服务器模式和手动复制 URL 模式。
"""
def __init__(
self,
client_id: str,
tenant_id: str = "common",
scopes: list[str] | None = None,
redirect_port: int = DEFAULT_CALLBACK_PORT,
redirect_path: str = DEFAULT_CALLBACK_PATH,
):
self._client_id = client_id
self._tenant_id = tenant_id
self._scopes = scopes or [
"https://graph.microsoft.com/User.Read",
"https://graph.microsoft.com/Chat.ReadWrite",
"https://graph.microsoft.com/ChannelMessage.Send",
"offline_access",
]
self._redirect_uri = f"http://localhost:{redirect_port}{redirect_path}"
self._code_verifier: str = ""
self._code_challenge: str = ""
self._state: str = ""
def build_authorize_url(self) -> str:
self._code_verifier = _generate_pkce_code_verifier()
self._code_challenge = _generate_pkce_code_challenge(self._code_verifier)
self._state = _generate_state()
params = {
"client_id": self._client_id,
"response_type": "code",
"redirect_uri": self._redirect_uri,
"scope": " ".join(self._scopes),
"state": self._state,
"code_challenge": self._code_challenge,
"code_challenge_method": "S256",
"response_mode": "query",
}
return f"{AUTHORIZE_URL.format(tenant=self._tenant_id)}?{urlencode(params)}"
async def execute_auto(self) -> dict[str, Any]:
authorize_url = self.build_authorize_url()
_CallbackHandler.callback_result = {}
server = HTTPServer(("localhost", DEFAULT_CALLBACK_PORT), _CallbackHandler)
server.timeout = 5
logger.info("MSTeams OAuth: opening browser for authorization...")
webbrowser.open(authorize_url)
max_attempts = 60
for _ in range(max_attempts):
server.handle_request()
if _CallbackHandler.callback_result:
break
await asyncio.sleep(1)
server.server_close()
result = dict(_CallbackHandler.callback_result)
if not result:
return {"error": "timeout", "error_description": "Authorization timed out"}
if "error" in result:
return result
if result.get("state") != self._state:
return {"error": "state_mismatch", "error_description": "State parameter mismatch"}
return await self._exchange_code(result["code"])
async def execute_manual(self) -> dict[str, Any]:
authorize_url = self.build_authorize_url()
print("\n" + "=" * 60)
print("Microsoft Teams OAuth 授权")
print("=" * 60)
print("\n请复制以下 URL 并在浏览器中打开:\n")
print(authorize_url)
print("\n授权完成后,将浏览器重定向到的完整 URL 粘贴到这里:")
print("=" * 60)
redirect_url = await asyncio.get_event_loop().run_in_executor(None, lambda: input("\n> ").strip())
if not redirect_url:
return {"error": "cancelled", "error_description": "No URL provided"}
parsed = urlparse(redirect_url)
qs = parse_qs(parsed.query)
code = qs.get("code", [""])[0]
state = qs.get("state", [""])[0]
error = qs.get("error", [""])[0]
error_desc = qs.get("error_description", [""])[0]
if error:
return {"error": error, "error_description": error_desc}
if not code:
return {"error": "no_code", "error_description": "No authorization code in URL"}
if state != self._state:
return {"error": "state_mismatch", "error_description": "State parameter mismatch"}
return await self._exchange_code(code)
async def execute(self) -> dict[str, Any]:
if needs_manual_oauth():
return await self.execute_manual()
return await self.execute_auto()
async def _exchange_code(self, code: str) -> dict[str, Any]:
token_endpoint = TOKEN_URL.format(tenant=self._tenant_id)
data = {
"client_id": self._client_id,
"grant_type": "authorization_code",
"code": code,
"redirect_uri": self._redirect_uri,
"code_verifier": self._code_verifier,
"scope": " ".join(self._scopes),
}
try:
async with aiohttp.ClientSession() as session:
async with session.post(token_endpoint, data=data) as resp:
if resp.status == 200:
result = await resp.json()
return {
"success": True,
"access_token": result.get("access_token", ""),
"refresh_token": result.get("refresh_token", ""),
"expires_in": result.get("expires_in", 3600),
"scope": result.get("scope", ""),
"token_type": result.get("token_type", "Bearer"),
}
body = await resp.text()
logger.warning(f"OAuth token exchange failed: HTTP {resp.status} - {body[:300]}")
return {"error": "token_exchange_failed", "error_description": body[:300]}
except Exception as e:
logger.error(f"OAuth token exchange error: {e}")
return {"error": "network_error", "error_description": str(e)}
def persist_tokens(
self,
result: dict[str, Any],
user_id: str,
storage_dir: str | None = None,
) -> bool:
if not result.get("success"):
return False
try:
from .credentials import DelegatedAuthStore
store = DelegatedAuthStore(storage_dir)
store.store_token(
user_id=user_id,
access_token=result["access_token"],
refresh_token=result["refresh_token"],
expires_in=result.get("expires_in", 3600),
)
logger.info(f"MSTeams OAuth: tokens persisted for user {user_id}")
return True
except Exception as e:
logger.error(f"MSTeams OAuth: failed to persist tokens: {e}")
return False