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, }