396 lines
15 KiB
Python
396 lines
15 KiB
Python
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,
|
|
} |