eight-hourr/backend/services/pts_client.py

396 lines
15 KiB
Python
Raw Permalink Normal View History

2026-07-17 06:55:36 +00:00
from __future__ import annotations
import base64
import json
import logging
from typing import Any
import httpx
from config import settings
from services.secret_box import decrypt_secret
from services.state_store import get_pts_session, set_pts_session
logger = logging.getLogger(__name__)
TASK_NAME_ALIASES = {
"development": "Implement",
"dev": "Implement",
}
_SESSION_HELP = (
"PTS session not found. Enter PTS account/password on the web UI, "
"or sync via Chrome extension."
)
def normalize_pts_base_url(url: str) -> str:
"""PTS REST APIs are hosted under /PTS, not the site root."""
base = url.rstrip("/")
if base.lower().endswith("/pts"):
return base
return f"{base}/PTS"
def _normalize_session_payload(data: Any) -> dict[str, Any]:
if not isinstance(data, dict):
raise ValueError("PTS login returned unexpected payload")
access = data.get("accessToken") or data.get("AccessToken")
if not access:
raise ValueError("PTS login did not return accessToken")
session = dict(data)
session["accessToken"] = access
if data.get("refreshToken") or data.get("RefreshToken"):
session["refreshToken"] = data.get("refreshToken") or data.get("RefreshToken")
group = data.get("groupCode") or data.get("GroupCode")
if not group and session.get("accessToken") and session["accessToken"].count(".") >= 2:
try:
payload = session["accessToken"].split(".")[1]
payload += "=" * (-len(payload) % 4)
decoded = json.loads(base64.urlsafe_b64decode(payload))
group = decoded.get("groupCode")
except (json.JSONDecodeError, ValueError, TypeError):
group = None
if group:
session["groupCode"] = str(group)
return session
class PTSClient:
def __init__(self) -> None:
self.base_url = normalize_pts_base_url(settings.pts_url)
def _require_session(self) -> dict[str, Any]:
session = get_pts_session()
if not session or not session.get("accessToken"):
raise ValueError(_SESSION_HELP)
return session
def _auth_headers(self) -> dict[str, str]:
session = self._require_session()
return {
"Authorization": f"Bearer {session['accessToken']}",
"Content-Type": "application/json",
}
def _group_code_from_session(self, session: dict[str, Any]) -> str | None:
group_code = session.get("groupCode")
if group_code:
return str(group_code)
refresh_token = session.get("refreshToken")
if not refresh_token or refresh_token.count(".") < 2:
return None
try:
payload = refresh_token.split(".")[1]
payload += "=" * (-len(payload) % 4)
decoded = json.loads(base64.urlsafe_b64decode(payload))
group_code = decoded.get("groupCode")
return str(group_code) if group_code else None
except (json.JSONDecodeError, ValueError, TypeError):
return None
def _stored_credentials(self) -> tuple[str, str] | None:
username = (getattr(settings, "pts_username", "") or "").strip()
encrypted = getattr(settings, "pts_password", "") or ""
if not username or not encrypted:
return None
try:
password = decrypt_secret(encrypted)
except ValueError:
return None
if not password:
return None
return username, password
async def login(
self,
username: str,
password: str,
*,
group_code: str | None = None,
) -> dict[str, Any]:
username = username.strip()
if not username or not password:
raise ValueError("PTS username and password are required")
errors: list[str] = []
session: dict[str, Any] | None = None
try:
session = await self._login_windows_ntlm(username, password)
except Exception as exc: # noqa: BLE001
errors.append(f"Windows/NTLM: {exc}")
logger.info("PTS Windows auth failed: %s", exc)
if session is None:
try:
session = await self._login_forms(username, password)
except Exception as exc: # noqa: BLE001
errors.append(f"Forms: {exc}")
logger.info("PTS Forms auth failed: %s", exc)
if session is None:
detail = "; ".join(errors) if errors else "unknown error"
raise ValueError(
"PTS login failed. Check account/password (Windows domain format "
f"DOMAIN\\user is supported). Details: {detail}"
)
if group_code:
session = await self.switch_group(session, group_code)
set_pts_session(session)
logger.info("PTS login succeeded for %s", username)
return {
"ok": True,
"message": "PTS login succeeded",
"group_code": session.get("groupCode"),
"has_session": True,
}
async def login_with_stored_credentials(self) -> dict[str, Any]:
creds = self._stored_credentials()
if not creds:
raise ValueError(
"No stored PTS credentials. Enter username/password on the web UI."
)
username, password = creds
return await self.login(username, password)
async def _login_windows_ntlm(self, username: str, password: str) -> dict[str, Any]:
try:
from httpx_ntlm import HttpNtlmAuth
except ImportError as exc:
raise ValueError(
"httpx-ntlm is not installed. Run: pip install httpx-ntlm"
) from exc
url = f"{self.base_url}/api/Login/WindowsAuthentication"
auth = HttpNtlmAuth(username, password)
async with httpx.AsyncClient(timeout=60.0, auth=auth) as client:
response = await client.get(url)
if response.status_code in {401, 403}:
raise ValueError("invalid username/password or NTLM denied")
response.raise_for_status()
data = response.json()
return _normalize_session_payload(data)
async def _login_forms(self, username: str, password: str) -> dict[str, Any]:
url = f"{self.base_url}/api/Login/FormsAuthentication"
bodies = [
{"userName": username, "password": password},
{"username": username, "password": password},
{"UserName": username, "Password": password},
{"account": username, "password": password},
{"id": username, "password": password},
{"employeeId": username, "password": password},
]
async with httpx.AsyncClient(timeout=60.0) as client:
last_status = None
for body in bodies:
response = await client.post(
url,
headers={"Content-Type": "application/json"},
json=body,
)
last_status = response.status_code
if response.status_code == 404:
break
if response.status_code >= 400:
continue
try:
data = response.json()
except json.JSONDecodeError:
continue
if isinstance(data, dict) and (data.get("accessToken") or data.get("AccessToken")):
return _normalize_session_payload(data)
raise ValueError(f"FormsAuthentication unavailable (HTTP {last_status})")
async def switch_group(self, session: dict[str, Any], group_code: str) -> dict[str, Any]:
url = f"{self.base_url}/api/Login/SwitchGroup"
async with httpx.AsyncClient(timeout=60.0) as client:
response = await client.post(
url,
headers={
"Authorization": f"Bearer {session['accessToken']}",
"Content-Type": "application/json",
},
json={"groupCode": group_code, "accessToken": session["accessToken"]},
)
response.raise_for_status()
data = response.json()
if not data:
session["groupCode"] = group_code
return session
return _normalize_session_payload(data)
async def refresh_access_token(self) -> str:
session = self._require_session()
refresh_token = session.get("refreshToken")
group_code = self._group_code_from_session(session)
if not refresh_token or not group_code:
return await self._relogin_or_raise(
"PTS session expired and cannot refresh. Re-login with password or Extension."
)
url = f"{self.base_url}/api/Login/RefreshToken"
async with httpx.AsyncClient(timeout=60.0) as client:
response = await client.post(
url,
headers={"Content-Type": "application/json"},
json={"refreshToken": refresh_token, "groupCode": group_code},
)
if response.status_code == 401:
return await self._relogin_or_raise(
"PTS session expired. Re-login with password or Extension."
)
response.raise_for_status()
data = response.json()
access_token = data.get("accessToken")
if not access_token:
raise ValueError("PTS refresh did not return accessToken")
session["accessToken"] = access_token
session["groupCode"] = group_code
set_pts_session(session)
logger.info("PTS access token refreshed")
return access_token
async def _relogin_or_raise(self, message: str) -> str:
if self._stored_credentials():
await self.login_with_stored_credentials()
session = get_pts_session() or {}
token = session.get("accessToken")
if token:
return str(token)
raise ValueError(message)
async def _request(
self,
method: str,
path: str,
*,
json_body: Any = None,
retry_on_unauthorized: bool = True,
) -> Any:
url = f"{normalize_pts_base_url(self.base_url)}{path}"
async with httpx.AsyncClient(timeout=60.0) as client:
for attempt in range(2):
response = await client.request(
method,
url,
headers=self._auth_headers(),
json=json_body,
)
if (
response.status_code == 401
and retry_on_unauthorized
and attempt == 0
and path != "/api/Login/RefreshToken"
):
await self.refresh_access_token()
continue
if response.status_code == 401:
raise ValueError(
"PTS session expired. Re-login with password or Extension."
)
if response.status_code == 404 and "/PTS/" not in str(response.request.url):
raise ValueError(
"PTS API returned 404 because the URL is missing /PTS. "
f"Set PTS URL to {normalize_pts_base_url(settings.pts_url)} "
"and run: make restart"
) from None
response.raise_for_status()
content_type = response.headers.get("content-type", "")
if "application/json" in content_type:
return response.json()
return response.text
raise RuntimeError("PTS request retry loop exhausted")
async def get_project_options(self) -> list[dict[str, Any]]:
result = await self._request("GET", "/api/ProjectCode/GetProjectOptionsForPersonal")
return result or []
async def get_task_type_options(self) -> list[dict[str, Any]]:
result = await self._request("GET", "/api/TaskType/GetTaskTypeOptions")
return result or []
async def search_reports(
self,
begin_date: str,
end_date: str,
*,
project_code: str | None = None,
task_id: int | None = None,
) -> list[dict[str, Any]]:
payload: dict[str, Any] = {
"beginDate": begin_date,
"endDate": end_date,
}
if project_code:
payload["projectCode"] = project_code
if task_id is not None:
payload["taskId"] = task_id
result = await self._request("POST", "/api/Report/SearchForMember", json_body=payload)
return result or []
async def create_report(self, payload: dict[str, Any]) -> int | str | None:
result = await self._request("POST", "/api/Report/CreateReport", json_body=payload)
if isinstance(result, dict):
return result.get("id")
return result
async def resolve_project_code(self, project_name: str | None = None) -> str:
name = project_name or settings.pts_project_name
options = await self.get_project_options()
for opt in options:
if opt.get("name") == name or opt.get("code") == name:
return opt["code"]
lowered = name.lower()
for opt in options:
if lowered in str(opt.get("name", "")).lower():
return opt["code"]
available = [f"{o.get('name')} ({o.get('code')})" for o in options[:20]]
raise ValueError(
f"Project '{name}' not found. Available: {', '.join(available)}"
)
async def resolve_task_id(self, task_name: str | None = None) -> int:
name = task_name or settings.pts_default_task_name
alias = TASK_NAME_ALIASES.get(name.lower())
if alias:
name = alias
options = await self.get_task_type_options()
for opt in options:
if opt.get("name") == name:
return int(opt["id"])
lowered = name.lower()
for opt in options:
if lowered in str(opt.get("name", "")).lower():
return int(opt["id"])
available = [o.get("name") for o in options if not o.get("disable")]
raise ValueError(
f"Task '{task_name}' not found. Available: {', '.join(available)}"
)
async def verify_connection(self) -> dict[str, Any]:
projects = await self.get_project_options()
tasks = await self.get_task_type_options()
project_code = await self.resolve_project_code()
task_id = await self.resolve_task_id()
task_name = next(
(t.get("name") for t in tasks if int(t.get("id", -1)) == task_id),
str(task_id),
)
return {
"base_url": self.base_url,
"project_count": len(projects),
"task_type_count": len(tasks),
"resolved_project": project_code,
"resolved_task_id": task_id,
"resolved_task_name": task_name,
}