-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdb.py
More file actions
611 lines (526 loc) · 23.7 KB
/
Copy pathdb.py
File metadata and controls
611 lines (526 loc) · 23.7 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
#!/usr/bin/env python3
import asyncio
import datetime
import hashlib
import json
import urllib.parse
from typing import Any
from typing import Literal
import asyncpg.connect_utils
import fastapi_structured_logging
from pydantic import BaseModel
from config import DefaultConfig
from config import config
from gitlab_model import GLEmojiAttributes
from gitlab_model import MergeRequestPayload
from gitlab_model import PipelinePayload
log = fastapi_structured_logging.get_logger()
__all__ = ["database", "dbh"]
class NoResetConnection(asyncpg.connection.Connection):
def __init__(
self,
protocol: asyncpg.protocol.protocol.BaseProtocol,
transport: object,
loop: asyncio.AbstractEventLoop,
addr: tuple[str, int] | str,
config: asyncpg.connect_utils._ClientConfiguration,
params: asyncpg.connect_utils._ConnectionParameters,
) -> None:
super().__init__(protocol, transport, loop, addr, config, params)
self._reset_query: list[str] = []
class DatabaseLifecycleHandler:
def __init__(self, conf: DefaultConfig):
self._pool: asyncpg.Pool | None = None
self._config = conf
async def connect(self):
log.debug("creating database connection pool")
self._pool = await asyncpg.create_pool(
dsn=self._config.DATABASE_URL,
server_settings={
"application_name": "notiteams-gitlab-mr-api",
},
connection_class=NoResetConnection,
init=self.init_connection,
min_size=self._config.DATABASE_POOL_MIN_SIZE,
max_size=self._config.DATABASE_POOL_MAX_SIZE,
)
# Simple check at startup, will validate database resolution and creds
async with await self.acquire() as connection:
await connection.fetchval("SELECT 1")
async def init_connection(self, conn: asyncpg.Connection) -> None:
log.debug("connecting to database")
await conn.set_type_codec("jsonb", encoder=json.dumps, decoder=json.loads, schema="pg_catalog")
if self._config.log_queries:
def relog(value: asyncpg.connection.LoggedQuery):
log.debug(
json.dumps(
{
"query": value.query,
"args": value.args,
"timeout": value.timeout,
"elapsed": value.elapsed,
"exception": str(value.exception),
},
default=str,
)
)
conn.add_query_logger(relog)
async def disconnect(self):
if self._pool:
await self._pool.close()
async def acquire(self) -> asyncpg.pool.PoolAcquireContext:
assert self._pool is not None
return self._pool.acquire()
class GitlabUser(BaseModel):
id: int
name: str
username: str
class GitlabApprovals(GitlabUser):
status: str
class EmojiEntry(BaseModel, extra="allow"):
object_kind: Literal["emoji"]
event_type: Literal["award"] | Literal["revoke"]
object_attributes: GLEmojiAttributes
user: GitlabUser
class DiscussionStats(BaseModel):
threads_total: int = 0
threads_resolved: int = 0
threads_unresolved: int = 0
comments_total: int = 0
comments_resolved: int = 0
comments_unresolved: int = 0
class MergeRequestExtraState(BaseModel):
version: int
opener: GitlabUser
approvers: dict[str, GitlabApprovals]
pipeline_statuses: dict[str, PipelinePayload]
emojis: dict[str, EmojiEntry]
discussion_stats: DiscussionStats | None = None
def has_unresolved_threads(extra_state: MergeRequestExtraState) -> bool:
"""Check if extra_state has unresolved threads."""
return extra_state.discussion_stats is not None and extra_state.discussion_stats.threads_unresolved > 0
class MergeRequestInfos(BaseModel):
merge_request_ref_id: int
merge_request_payload: MergeRequestPayload
merge_request_extra_state: MergeRequestExtraState
head_pipeline_id: int | None
def compute_mri_fingerprint(mri: MergeRequestInfos) -> str:
"""Compute a stable fingerprint from MRI data for deduplication."""
datasource = {
"mri_payload": mri.merge_request_payload.model_dump(),
"mri_extra_state": mri.merge_request_extra_state.model_dump(),
"head_pipeline_id": mri.head_pipeline_id,
}
return hashlib.sha256(json.dumps(datasource, sort_keys=True, default=str).encode()).hexdigest()
def make_mr_summary(mri: MergeRequestInfos) -> str:
"""Create a summary string for Teams message fallback."""
return (
f"MR {mri.merge_request_payload.object_attributes.state}:"
f" {mri.merge_request_payload.object_attributes.title}\n"
f"on {mri.merge_request_payload.project.path_with_namespace}"
)
class DBHelper:
def __init__(self, database: DatabaseLifecycleHandler):
self.db: DatabaseLifecycleHandler = database
async def get_gitlab_instance_id_from_url(self, urlstr: str) -> int:
url = urllib.parse.urlparse(urlstr)
if not url.netloc:
raise ValueError(f"unable to determine gitlab host from {urlstr}")
hostname = url.netloc.lower()
gli_id = await self._generic_norm_upsert(
table="gitlab_instance",
identity_col="gitlab_instance_id",
select_attrs={"hostname": hostname},
)
assert isinstance(gli_id, int)
return gli_id
async def get_or_create_merge_request_ref_id(self, merge_request: MergeRequestPayload) -> int:
"""
Get or create an MR ref without updating payload.
Used for OOO check - we need the ID to query message refs,
but we don't want to corrupt the payload with stale data.
Only sets initial state on INSERT, never updates existing records.
"""
gitlab_instance_id = await self.get_gitlab_instance_id_from_url(merge_request.object_attributes.url)
merge_ref_id = await self._generic_norm_upsert(
table="merge_request_ref",
identity_col="merge_request_ref_id",
select_attrs={
"gitlab_instance_id": gitlab_instance_id,
"gitlab_project_id": merge_request.object_attributes.target_project_id,
"gitlab_merge_request_iid": merge_request.object_attributes.iid,
},
insert_only_vals={
"gitlab_merge_request_id": merge_request.object_attributes.id,
"head_pipeline_id": merge_request.object_attributes.head_pipeline_id,
"merge_request_payload": merge_request.model_dump(),
"merge_request_extra_state": {
"version": 1,
"opener": {
"id": merge_request.user.id,
"name": merge_request.user.name,
"username": merge_request.user.username,
},
"approvers": {},
"pipeline_statuses": {},
"emojis": {},
},
},
)
assert isinstance(merge_ref_id, int)
return merge_ref_id
async def update_merge_request_ref_payload(
self, merge_request_ref_id: int, merge_request: MergeRequestPayload
) -> MergeRequestInfos:
"""
Update MR ref payload after OOO check has passed.
Called only for non-OOO events to update the stored payload.
"""
connection: asyncpg.Connection
async with await database.acquire() as connection:
row = await connection.fetchrow(
"""UPDATE merge_request_ref
SET gitlab_merge_request_id = $1,
head_pipeline_id = $2,
merge_request_payload = $3
WHERE merge_request_ref_id = $4
RETURNING merge_request_ref_id, merge_request_payload,
merge_request_extra_state, head_pipeline_id""",
merge_request.object_attributes.id,
merge_request.object_attributes.head_pipeline_id,
merge_request.model_dump(),
merge_request_ref_id,
)
assert row is not None
return MergeRequestInfos(**row)
async def get_merge_request_ref_infos(self, merge_request: MergeRequestPayload) -> MergeRequestInfos:
"""
Get or create MR ref AND update payload (legacy behavior).
Note: This updates payload on every call. For OOO-safe behavior,
use get_or_create_merge_request_ref_id() + update_merge_request_ref_payload().
"""
gitlab_instance_id = await self.get_gitlab_instance_id_from_url(merge_request.object_attributes.url)
merge_ref = await self._generic_norm_upsert(
table="merge_request_ref",
identity_col="merge_request_ref_id",
select_attrs={
"gitlab_instance_id": gitlab_instance_id,
"gitlab_project_id": merge_request.object_attributes.target_project_id,
"gitlab_merge_request_iid": merge_request.object_attributes.iid,
},
extra_insert_and_update_vals={
"gitlab_merge_request_id": merge_request.object_attributes.id,
"head_pipeline_id": merge_request.object_attributes.head_pipeline_id,
"merge_request_payload": merge_request.model_dump(),
},
insert_only_vals={
"merge_request_extra_state": {
"version": 1,
"opener": {
"id": merge_request.user.id,
"name": merge_request.user.name,
"username": merge_request.user.username,
},
"approvers": {},
"pipeline_statuses": {},
"emojis": {},
},
},
extra_sel_cols=["merge_request_payload", "merge_request_extra_state", "head_pipeline_id"],
)
assert isinstance(merge_ref, asyncpg.Record)
return MergeRequestInfos(**merge_ref)
async def update_discussion_stats(
self, merge_request_ref_id: int, stats: "DiscussionStats"
) -> MergeRequestExtraState:
"""Update discussion_stats in merge_request_extra_state."""
connection: asyncpg.Connection
async with await database.acquire() as connection:
row = await connection.fetchrow(
"""UPDATE merge_request_ref
SET merge_request_extra_state = jsonb_set(
merge_request_extra_state,
'{discussion_stats}',
$1::jsonb
)
WHERE merge_request_ref_id = $2
RETURNING merge_request_extra_state""",
stats.model_dump(),
merge_request_ref_id,
)
assert row is not None
return MergeRequestExtraState(**row["merge_request_extra_state"])
async def get_mri_from_url_pid_mriid(
self,
url: str,
project_id: int,
mr_iid: int,
) -> MergeRequestInfos | None:
gitlab_instance_id = await self.get_gitlab_instance_id_from_url(url)
connection: asyncpg.Connection
async with await database.acquire() as connection:
row = await connection.fetchrow(
"""SELECT
merge_request_ref_id,
merge_request_payload,
merge_request_extra_state,
head_pipeline_id
FROM merge_request_ref
WHERE
gitlab_instance_id = $1
AND gitlab_project_id = $2
AND gitlab_merge_request_iid = $3
""",
gitlab_instance_id,
project_id,
mr_iid,
)
if row is not None:
return MergeRequestInfos(**row)
return None
async def _generic_norm_upsert(
self,
*,
table: str,
identity_col: str,
select_attrs: dict[str, Any],
extra_insert_and_update_vals: dict[str, Any] | None = None,
insert_only_vals: dict[str, Any] | None = None,
extra_sel_cols: list[str] | None = None,
) -> Any:
if extra_insert_and_update_vals is None:
extra_insert_and_update_vals = {}
if insert_only_vals is None:
insert_only_vals = {}
if extra_sel_cols is None:
extra_sel_cols = []
sel_cols: list[str] = [f'"{identity_col}"']
sel_cols.extend([f'"{extra_col}"' for extra_col in extra_sel_cols])
sel_where: list[str] = []
sel_args: list[Any] = []
upd_set: list[str] = []
upd_where: list[str] = []
upd_args: list[Any] = []
for k, v in extra_insert_and_update_vals.items():
upd_args.append(v)
upd_set.append(f'"{k}" = ${len(upd_args)}')
for k, v in select_attrs.items():
sel_args.append(v)
sel_where.append(f'"{k}" = ${len(sel_args)}')
upd_args.append(v)
upd_where.append(f'"{k}" = ${len(upd_args)}')
if len(upd_set):
query = f"""UPDATE "{table}"
SET {", ".join(upd_set)}
WHERE {" AND ".join(upd_where)}
RETURNING {", ".join(sel_cols)}"""
args = upd_args
else:
query = f"""SELECT
{", ".join(sel_cols)}
FROM "{table}"
WHERE {" AND ".join(sel_where)}"""
args = sel_args
connection: asyncpg.Connection
async with await database.acquire() as connection:
row = await connection.fetchrow(
query,
*args,
)
if row is None:
try:
ins_col = []
ins_args = []
for k, v in select_attrs.items():
ins_col.append(f'"{k}"')
ins_args.append(v)
for k, v in extra_insert_and_update_vals.items():
ins_col.append(f'"{k}"')
ins_args.append(v)
for k, v in insert_only_vals.items():
ins_col.append(f'"{k}"')
ins_args.append(v)
row = await connection.fetchrow(
f"""
INSERT INTO "{table}" (
{", ".join(ins_col)}
) VALUES (
{", ".join(["$" + str(i + 1) for i in range(len(ins_col))])}
) RETURNING {", ".join(sel_cols)}
""",
*ins_args,
)
except asyncpg.exceptions.UniqueViolationError:
row = await connection.fetchrow(
query,
*args,
)
assert row is not None
if len(extra_sel_cols) == 0:
return row[identity_col]
return row
async def any_message_needs_update(self, merge_request_ref_id: int, payload_fingerprint: str) -> bool:
"""
Check if ANY message for an MR needs updating (pre-check for deduplication).
Returns True if at least one message hasn't been processed with the given fingerprint.
Used to skip expensive GitLab API calls when all messages are already up-to-date.
"""
connection: asyncpg.Connection
async with await database.acquire() as connection:
row = await connection.fetchrow(
"""SELECT EXISTS(
SELECT 1 FROM merge_request_message_ref
WHERE merge_request_ref_id = $1
AND message_id IS NOT NULL
AND (last_processed_fingerprint IS NULL
OR last_processed_fingerprint != $2)
) as needs_update""",
merge_request_ref_id,
payload_fingerprint,
)
return bool(row["needs_update"]) if row else False
async def upsert_pending_mr_refresh(
self,
merge_request_ref_id: int,
payload_type: str,
debounce_seconds: float = 2.0,
) -> bool:
"""
Insert or update a pending MR refresh entry.
Returns True if this is a new entry (first event), False if debounced (subsequent event).
On first event: immediate processing (process_after = now).
On subsequent: updates last_event_at and extends process_after by debounce_seconds.
"""
connection: asyncpg.Connection
async with await database.acquire() as connection:
result: str = await connection.execute(
"""INSERT INTO pending_mr_refresh
(merge_request_ref_id, payload_type, first_event_at, last_event_at, process_after)
VALUES ($1, $2, now(), now(), now())
ON CONFLICT (merge_request_ref_id) DO UPDATE
SET last_event_at = now(),
process_after = now() + $3::interval""",
merge_request_ref_id,
payload_type,
datetime.timedelta(seconds=debounce_seconds),
)
return result == "INSERT 0 1"
async def get_pending_refreshes(self, limit: int = 50) -> list[dict[str, Any]]:
"""Get pending refreshes ready for processing."""
connection: asyncpg.Connection
async with await database.acquire() as connection:
rows = await connection.fetch(
"""SELECT pmr.merge_request_ref_id, pmr.payload_type,
pmr.first_event_at, pmr.last_event_at,
mr.merge_request_payload, mr.merge_request_extra_state,
mr.head_pipeline_id
FROM pending_mr_refresh pmr
JOIN merge_request_ref mr USING (merge_request_ref_id)
WHERE pmr.process_after <= now()
ORDER BY pmr.process_after
LIMIT $1
FOR UPDATE OF pmr SKIP LOCKED""",
limit,
)
return [dict(row) for row in rows]
async def refresh_mr_payload_from_api(
self, merge_request_ref_id: int, api_data: dict[str, Any]
) -> MergeRequestInfos:
"""Update stored MR payload with fresh state from GitLab API.
Syncs state, title, draft, merge status, branches, and pipeline ID.
Assignees/reviewers are not updated (API response lacks email field required by GLUser).
"""
connection: asyncpg.Connection
async with await database.acquire() as connection:
async with connection.transaction():
row = await connection.fetchrow(
"""SELECT merge_request_ref_id, merge_request_payload,
merge_request_extra_state, head_pipeline_id
FROM merge_request_ref
WHERE merge_request_ref_id = $1
FOR UPDATE""",
merge_request_ref_id,
)
assert row is not None
payload = row["merge_request_payload"]
oa = payload.get("object_attributes", {})
for field in (
"state",
"title",
"draft",
"detailed_merge_status",
"source_branch",
"target_branch",
):
if field in api_data:
oa[field] = api_data[field]
# GitLab REST API returns updated_at as ISO 8601 ("...Z");
# webhook payloads use "YYYY-MM-DD HH:MM:SS UTC". Normalize
# to webhook format so downstream fromisoformat parsing
# (with the " UTC" -> "+00:00" replace) keeps working.
# Defensive: if GitLab ever returns a naive datetime (no
# offset), assume UTC rather than letting astimezone() apply
# the host's local TZ. If parsing fails outright, log and
# keep the stored value — never raise here, otherwise the
# whole pending_mr_refresh row gets stuck retrying forever.
if "updated_at" in api_data and api_data["updated_at"]:
raw = api_data["updated_at"]
try:
parsed = datetime.datetime.fromisoformat(raw.replace("Z", "+00:00"))
if parsed.tzinfo is None:
parsed = parsed.replace(tzinfo=datetime.UTC)
oa["updated_at"] = parsed.astimezone(datetime.UTC).strftime("%Y-%m-%d %H:%M:%S UTC")
except (ValueError, TypeError) as exc:
log.warning(
"could not parse api updated_at, keeping stored value",
merge_request_ref_id=merge_request_ref_id,
raw=raw,
error=str(exc),
)
if "draft" in api_data:
oa["work_in_progress"] = api_data["draft"]
# Synthesize `action` from state so cards/render.py picks the
# right icon (CodeTextOff for close, Merge for merge). The
# renderer keys off action, not state, so we MUST set it.
# Only mutate on terminal/reopen transitions; otherwise keep
# the webhook-recorded action to avoid masking real events.
api_state = api_data.get("state")
if api_state == "merged" and oa.get("action") != "merge":
oa["action"] = "merge"
elif api_state == "closed" and oa.get("action") != "close":
oa["action"] = "close"
elif api_state == "opened" and oa.get("action") in ("close", "merge"):
oa["action"] = "reopen"
head_pipeline_id = row["head_pipeline_id"]
api_pipeline = api_data.get("head_pipeline")
if api_pipeline and api_pipeline.get("id"):
oa["head_pipeline_id"] = api_pipeline["id"]
head_pipeline_id = api_pipeline["id"]
payload["object_attributes"] = oa
await connection.execute(
"""UPDATE merge_request_ref
SET merge_request_payload = $1, head_pipeline_id = $2
WHERE merge_request_ref_id = $3""",
payload,
head_pipeline_id,
merge_request_ref_id,
)
# extra_state is read from the pre-update `row` snapshot. Safe today
# because this function does not mutate extra_state; if that ever
# changes, re-read it after the UPDATE or RETURNING it.
return MergeRequestInfos(
merge_request_ref_id=merge_request_ref_id,
merge_request_payload=payload,
merge_request_extra_state=row["merge_request_extra_state"],
head_pipeline_id=head_pipeline_id,
)
async def delete_pending_refresh(self, merge_request_ref_id: int) -> None:
"""Delete a pending refresh after processing."""
connection: asyncpg.Connection
async with await database.acquire() as connection:
await connection.execute(
"DELETE FROM pending_mr_refresh WHERE merge_request_ref_id = $1",
merge_request_ref_id,
)
database = DatabaseLifecycleHandler(config)
dbh = DBHelper(database)