Skip to content

Commit db6b1c6

Browse files
author
testgen-ci-bot
committed
Merge remote-tracking branch 'origin/enterprise' into feat/TG-1115-add-job-execution-relationship-to-run-models
2 parents 7ff7b60 + b3324a3 commit db6b1c6

13 files changed

Lines changed: 219 additions & 59 deletions

File tree

testgen/api/deps.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@
66
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
77
from sqlalchemy import select
88

9-
from testgen.common.auth import authorize_token, decode_jwt_token
9+
from testgen.common.auth import AuthError, authorize_token, decode_jwt_token
1010
from testgen.common.enums import PublicJobKey
1111
from testgen.common.models import Session, _current_session_wrapper, get_current_session
1212
from testgen.common.models.job_execution import JobExecution
@@ -48,7 +48,7 @@ def get_authorized_user(credentials: HTTPAuthorizationCredentials = _bearer_secu
4848

4949
try:
5050
payload = decode_jwt_token(credentials.credentials)
51-
except ValueError:
51+
except AuthError:
5252
raise _invalid from None
5353

5454
username = payload.get("username")
@@ -58,7 +58,7 @@ def get_authorized_user(credentials: HTTPAuthorizationCredentials = _bearer_secu
5858
session = get_current_session()
5959
try:
6060
return authorize_token(credentials.credentials, username, session)
61-
except ValueError:
61+
except AuthError:
6262
raise _invalid from None
6363

6464

testgen/common/auth.py

Lines changed: 8 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,10 @@
1010
LOG = logging.getLogger("testgen")
1111

1212

13+
class AuthError(Exception):
14+
"""Token authentication failed — invalid/expired token, unknown user, or revoked token."""
15+
16+
1317
def get_jwt_signing_key() -> bytes:
1418
"""Decode the base64-encoded JWT signing key from settings."""
1519
return base64.b64decode(settings.JWT_HASHING_KEY_B64.encode("ascii"))
@@ -27,13 +31,13 @@ def create_jwt_token(username: str, expiry_seconds: int = 86400) -> str:
2731
def decode_jwt_token(token_str: str) -> dict:
2832
"""Decode and validate a JWT token. Returns the payload dict.
2933
30-
Raises ValueError if the token is invalid or expired.
34+
Raises ``AuthError`` if the token is invalid or expired.
3135
PyJWT auto-validates the standard ``exp`` claim during decode.
3236
"""
3337
try:
3438
return jwt.decode(token_str, get_jwt_signing_key(), algorithms=["HS256"])
3539
except jwt.InvalidTokenError as e:
36-
raise ValueError(f"Invalid token: {e}") from e
40+
raise AuthError(f"Invalid token: {e}") from e
3741

3842

3943
def authorize_token(token_str: str, username: str, session):
@@ -48,11 +52,11 @@ def authorize_token(token_str: str, username: str, session):
4852

4953
user = session.scalars(select(User).where(func.lower(User.username) == func.lower(username))).first()
5054
if user is None:
51-
raise ValueError("User not found")
55+
raise AuthError("User not found")
5256

5357
token_record = session.scalars(select(OAuth2Token).where(OAuth2Token.access_token == token_str)).first()
5458
if token_record and token_record.access_token_revoked_at:
55-
raise ValueError("Token has been revoked")
59+
raise AuthError("Token has been revoked")
5660

5761
return user
5862

testgen/mcp/exceptions.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,14 @@ class MCPUserError(Exception):
2020
"""
2121

2222

23+
class MCPAuthenticationError(MCPUserError):
24+
"""Token authenticated at the transport but failed authorization at the tool boundary.
25+
26+
Raised when the token's user no longer exists or the token was revoked — distinct from
27+
per-project permission denial. Carries a uniform, re-authenticate message for the client.
28+
"""
29+
30+
2331
class MCPPermissionDenied(MCPUserError):
2432
"""Raised when access is denied due to insufficient project permissions."""
2533

testgen/mcp/permissions.py

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -2,14 +2,17 @@
22

33
import contextvars
44
import functools
5+
import logging
56
from collections.abc import Callable
67
from dataclasses import dataclass
78

89
from testgen.common.models.project_membership import ProjectMembership
910
from testgen.common.models.user import User
10-
from testgen.mcp.exceptions import MCPPermissionDenied
11+
from testgen.mcp.exceptions import MCPAuthenticationError, MCPPermissionDenied
1112
from testgen.utils.plugins import PluginHook
1213

14+
LOG = logging.getLogger("testgen")
15+
1316
_NOT_SET = object()
1417

1518
_mcp_username: contextvars.ContextVar[str | None] = contextvars.ContextVar("mcp_username", default=None)
@@ -83,7 +86,7 @@ def get_authorized_mcp_user() -> User:
8386
Checks user existence and token revocation status.
8487
Must be called within @with_database_session scope.
8588
"""
86-
from testgen.common.auth import authorize_token
89+
from testgen.common.auth import AuthError, authorize_token
8790
from testgen.common.models import get_current_session
8891

8992
username = _mcp_username.get()
@@ -92,7 +95,13 @@ def get_authorized_mcp_user() -> User:
9295

9396
token_str = _mcp_token.get()
9497
session = get_current_session()
95-
return authorize_token(token_str or "", username, session)
98+
try:
99+
return authorize_token(token_str or "", username, session)
100+
except AuthError as err:
101+
LOG.warning("MCP token authorization failed: %s", err)
102+
raise MCPAuthenticationError(
103+
"Authentication failed: your access token is no longer valid. Please sign in again."
104+
) from err
96105

97106

98107
def _compute_project_permissions(user: User, permission: str) -> ProjectPermissions:

testgen/mcp/server.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@
99
from starlette.applications import Starlette
1010

1111
from testgen import settings
12-
from testgen.common.auth import decode_jwt_token
12+
from testgen.common.auth import AuthError, decode_jwt_token
1313
from testgen.mcp.permissions import set_mcp_token, set_mcp_username
1414

1515
LOG = logging.getLogger("testgen")
@@ -65,7 +65,7 @@ async def verify_token(self, token: str) -> AccessToken | None:
6565
scopes=[],
6666
expires_at=int(payload["exp"]),
6767
)
68-
except (ValueError, KeyError):
68+
except (AuthError, KeyError):
6969
return None
7070

7171

testgen/mcp/tools/common.py

Lines changed: 62 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -300,6 +300,22 @@ def parse_run_status_filter(value: str) -> list[JobStatus]:
300300
return statuses
301301

302302

303+
class FailureGroupBy(StrEnum):
304+
"""User-facing values accepted for the ``group_by`` argument on ``get_failure_summary``."""
305+
306+
TEST_TYPE = "test_type"
307+
TABLE = "table"
308+
COLUMN = "column"
309+
310+
311+
def parse_failure_group_by(value: str) -> FailureGroupBy:
312+
try:
313+
return FailureGroupBy(value)
314+
except ValueError as err:
315+
valid = ", ".join(g.value for g in FailureGroupBy)
316+
raise MCPUserError(f"Invalid group_by `{value}`. Valid values: {valid}") from err
317+
318+
303319
def format_run_duration(started_at: datetime | None, completed_at: datetime | None) -> str | None:
304320
"""Render an elapsed duration as ``Xs`` / ``Xm Ys`` / ``Xh Ym``. Returns ``None`` if either bound is missing."""
305321
if not started_at or not completed_at:
@@ -661,6 +677,52 @@ def resolve_notification(notification_id: str) -> NotificationSettings:
661677
return notif
662678

663679

680+
def resolve_aggregate_scope(
681+
project_code: str | None,
682+
test_suite_id: str | None = None,
683+
table_group_id: str | None = None,
684+
) -> list[str]:
685+
"""Validate optional project / test-suite / table-group scope for a cross-run aggregation.
686+
687+
Resolves any supplied suite or table group (existence + access via the
688+
``resolve_*`` helpers), and — when ``project_code`` is also given — requires the
689+
resolved entity to belong to it, raising a clear error on a cross-project mismatch
690+
so the query never silently returns empty. Returns the project codes to scope the
691+
aggregation to.
692+
"""
693+
perms = get_project_permissions()
694+
if project_code:
695+
perms.verify_access(project_code, not_found=MCPResourceNotAccessible("Project", project_code))
696+
697+
scoped_projects: set[str] = set()
698+
if test_suite_id:
699+
suite = resolve_test_suite(test_suite_id)
700+
if project_code and suite.project_code != project_code:
701+
raise MCPUserError(
702+
f"Test suite `{test_suite_id}` belongs to project `{suite.project_code}`, not `{project_code}`."
703+
)
704+
scoped_projects.add(suite.project_code)
705+
if table_group_id:
706+
table_group = resolve_table_group(table_group_id)
707+
if project_code and table_group.project_code != project_code:
708+
raise MCPUserError(
709+
f"Table group `{table_group_id}` belongs to project `{table_group.project_code}`, not `{project_code}`."
710+
)
711+
scoped_projects.add(table_group.project_code)
712+
713+
if len(scoped_projects) > 1:
714+
# Suite and table group resolve to different projects (only reachable when no
715+
# project_code pins the scope). The two filters would AND to an empty result —
716+
# reject rather than silently return nothing.
717+
raise MCPUserError("The test suite and table group belong to different projects — narrow to one scope.")
718+
719+
if project_code:
720+
return [project_code]
721+
if scoped_projects:
722+
return list(scoped_projects)
723+
return perms.allowed_codes
724+
725+
664726
# Notification event-type labels.
665727

666728
class NotificationEventLabel(StrEnum):

testgen/mcp/tools/discovery.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
from testgen.common.models.project import Project
55
from testgen.common.models.test_run import TestRun
66
from testgen.common.models.test_suite import TestSuite
7+
from testgen.mcp.exceptions import MCPResourceNotAccessible
78
from testgen.mcp.permissions import get_project_permissions, mcp_permission
89
from testgen.mcp.tools.common import DocGroup, resolve_table_group, validate_limit, validate_page
910
from testgen.mcp.tools.markdown import MdDoc
@@ -63,7 +64,7 @@ def list_test_suites(project_code: str) -> str:
6364
return "Missing required parameter `project_code`."
6465

6566
perms = get_project_permissions()
66-
perms.verify_access(project_code, not_found=f"No test suites found for project `{project_code}`.")
67+
perms.verify_access(project_code, not_found=MCPResourceNotAccessible("Project", project_code))
6768

6869
summaries = TestSuite.select_summary(project_code)
6970

testgen/mcp/tools/test_results.py

Lines changed: 22 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -12,11 +12,14 @@
1212
from testgen.mcp.permissions import get_project_permissions, mcp_permission
1313
from testgen.mcp.tools.common import (
1414
DocGroup,
15+
FailureGroupBy,
1516
format_page_footer,
1617
format_page_info,
18+
parse_failure_group_by,
1719
parse_result_status,
1820
parse_since_arg,
1921
parse_uuid,
22+
resolve_aggregate_scope,
2023
resolve_test_type,
2124
validate_limit,
2225
validate_page,
@@ -27,6 +30,12 @@
2730

2831
_DEFAULT_SEARCH_STATUSES = [TestResultStatus.Failed, TestResultStatus.Warning]
2932

33+
_MODEL_GROUP_COLUMN = {
34+
FailureGroupBy.TEST_TYPE: "test_type",
35+
FailureGroupBy.TABLE: "table_name",
36+
FailureGroupBy.COLUMN: "column_names",
37+
}
38+
3039

3140
@with_database_session
3241
@mcp_permission("view")
@@ -167,20 +176,20 @@ def get_failure_summary(
167176
group_by: Group failures by 'test_type', 'table', or 'column' (default: 'test_type').
168177
"""
169178
perms = get_project_permissions()
179+
group = parse_failure_group_by(group_by)
170180

171181
if not any((job_execution_id, test_suite_id, since)):
172182
raise MCPUserError(
173183
"Provide 'job_execution_id' for a single run, or 'test_suite_id' or 'project_code' "
174184
"to aggregate across runs. 'since' is required when 'test_suite_id' is not provided."
175185
)
176-
if group_by in ("table", "column") and not (job_execution_id or test_suite_id):
186+
if group in (FailureGroupBy.TABLE, FailureGroupBy.COLUMN) and not (job_execution_id or test_suite_id):
177187
raise MCPUserError(
178-
f"'{group_by}' grouping requires a single-suite scope. "
188+
f"'{group}' grouping requires a single-suite scope. "
179189
"Provide 'job_execution_id' or 'test_suite_id'."
180190
)
181191

182-
model_group_map = {"table": "table_name", "column": "column_names"}
183-
model_group_by = model_group_map.get(group_by, group_by)
192+
model_group_by = _MODEL_GROUP_COLUMN[group]
184193

185194
scope_label: str
186195
test_run_id = None
@@ -197,15 +206,7 @@ def get_failure_summary(
197206
scope_label = f"run `{job_execution_id}`"
198207
project_codes = perms.allowed_codes
199208
else:
200-
if project_code:
201-
perms.verify_access(project_code, not_found=MCPResourceNotAccessible("Project", project_code))
202-
project_codes = [project_code]
203-
else:
204-
project_codes = perms.allowed_codes
205-
if test_suite_uuid is not None:
206-
suite = TestSuite.get_regular(test_suite_uuid)
207-
if suite is None or not perms.has_access(suite.project_code):
208-
raise MCPResourceNotAccessible("Test suite", test_suite_id)
209+
project_codes = resolve_aggregate_scope(project_code, test_suite_id=test_suite_id)
209210
scope_parts = []
210211
if project_code:
211212
scope_parts.append(f"project `{project_code}`")
@@ -227,22 +228,22 @@ def get_failure_summary(
227228
return f"No confirmed failures found for {scope_label}."
228229

229230
total = sum(row[-1] for row in failures)
230-
if group_by == "test_type":
231+
if group is FailureGroupBy.TEST_TYPE:
231232
type_names = {tt.test_type: tt.test_name_short for tt in TestType.select_where(TestType.active == "Y")}
232233

233234
doc = MdDoc()
234235
doc.heading(1, f"Failure Summary — {scope_label}")
235236
doc.text(f"**Total confirmed failures (Failed + Warning):** {total}")
236237

237-
if group_by == "test_type":
238+
if group is FailureGroupBy.TEST_TYPE:
238239
headers = ["Test Type", "Severity", "Count"]
239240
rows = []
240241
for row in failures:
241242
code, status, count = row[0], row[1], row[-1]
242243
name = type_names.get(code, code)
243244
severity = status.value if status else "Unknown"
244245
rows.append([name, severity, count])
245-
elif group_by == "column":
246+
elif group is FailureGroupBy.COLUMN:
246247
headers = ["Column", "Count"]
247248
rows = []
248249
for row in failures:
@@ -253,9 +254,9 @@ def get_failure_summary(
253254
headers = ["Table Name", "Count"]
254255
rows = [[row[0], row[-1]] for row in failures]
255256

256-
doc.table(headers, rows, code=[0] if group_by == "table" else None)
257+
doc.table(headers, rows, code=[0] if group is FailureGroupBy.TABLE else None)
257258

258-
if group_by == "test_type":
259+
if group is FailureGroupBy.TEST_TYPE:
259260
doc.text(
260261
"Check `testgen://test-types` to understand what each test type checks "
261262
"and `get_test_type(test_type='...')` to fetch more details."
@@ -450,12 +451,9 @@ def get_failure_trend(
450451
valid = ", ".join(v.value for v in BucketInterval)
451452
raise MCPUserError(f"Invalid `bucket`: `{bucket}`. Valid values: {valid}") from err
452453

453-
perms = get_project_permissions()
454-
if project_code:
455-
perms.verify_access(project_code, not_found=MCPResourceNotAccessible("Project", project_code))
456-
project_codes = [project_code]
457-
else:
458-
project_codes = perms.allowed_codes
454+
project_codes = resolve_aggregate_scope(
455+
project_code, test_suite_id=test_suite_id, table_group_id=table_group_id
456+
)
459457

460458
anchor_today = datetime.now(UTC).date()
461459
if exclude_today:

0 commit comments

Comments
 (0)