Mirrored from GitHub
github.com/roostorg/osprey
5.0 kB
126 lines
1from collections.abc import Generator
2from contextlib import contextmanager
3from typing import Any
4
5from osprey.engine.language_types.entities import EntityT
6from osprey.worker.lib.osprey_shared.labels import EntityLabels
7from osprey.worker.lib.osprey_shared.logging import get_logger
8from osprey.worker.lib.storage.labels import LabelsServiceBase
9from osprey.worker.lib.storage.postgres import Model, init_from_config, scoped_session
10from sqlalchemy import Column, String, select
11from sqlalchemy.dialects.postgresql import JSONB, insert
12
13logger = get_logger(__name__)
14
15
16class EntityLabelsModel(Model):
17 """SQLAlchemy model for storing entity labels in PostgreSQL"""
18
19 __tablename__ = 'entity_labels'
20
21 entity_key = Column(String, primary_key=True)
22 labels = Column(JSONB, nullable=False)
23
24 def __str__(self) -> str:
25 return f'EntityLabelsModel(entity_key={self.entity_key}, labels={self.labels})'
26
27
28class PostgresLabelsService(LabelsServiceBase):
29 """
30 PostgreSQL-backed implementation of LabelsServiceBase.
31
32 This service stores entity labels in a PostgreSQL database using SQLAlchemy.
33 It provides atomic read-modify-write operations through database transactions.
34 """
35
36 def __init__(self, database: str = 'osprey_db') -> None:
37 """
38 Initialize the PostgreSQL labels service.
39 Note: This will not init the postgres connection; To do that,
40 initialize() must be called (which is called by the LabelsProvider
41 by default)
42
43 Args:
44 database: The database name to use. Defaults to 'osprey_db'.
45 """
46 super().__init__()
47 self._database_name: str = database
48
49 def initialize(self) -> None:
50 init_from_config(self._database_name)
51 logger.info(f'Initialized PostgresLabelsService with database: {self._database_name}')
52
53 def read_labels(self, entity: EntityT[Any]) -> EntityLabels:
54 """
55 Read labels for an entity from PostgreSQL.
56
57 Returns an empty EntityLabels if the entity has no labels.
58 """
59 entity_key = str(entity)
60
61 with scoped_session(database=self._database_name) as session:
62 stmt = select(EntityLabelsModel).where(EntityLabelsModel.entity_key == entity_key)
63 result = session.scalars(stmt).first()
64
65 if result is None:
66 logger.debug(f'No labels found for entity {entity_key}')
67 return EntityLabels()
68
69 labels = EntityLabels.deserialize(result.labels)
70 logger.debug(f'Read labels for entity {entity_key}', result)
71 return labels
72
73 @contextmanager
74 def read_modify_write_labels_atomically(self, entity: EntityT[Any]) -> Generator[EntityLabels, None, None]:
75 """
76 Context manager for atomic read-modify-write operations.
77
78 This context manager:
79 1. Opens a database transaction
80 2. Acquires a row-level lock using SELECT FOR UPDATE
81 3. Reads and returns the current labels
82 4. Yields control to the caller (LabelsProvider)
83 5. The caller modifies the labels IN PLACE
84 6. On exit, writes the modified labels and commits the transaction
85
86 The key insight: The caller modifies the yielded labels object directly,
87 and this context manager persists those changes atomically.
88
89 For systems that don't need locking (e.g., in-memory stores), this can
90 be simplified to:
91 ```py
92 labels = self.read_labels(entity)
93 yield labels
94 # write the labels here
95 """
96 entity_key = str(entity)
97
98 with scoped_session(commit=False, database=self._database_name) as session:
99 try:
100 # Use SELECT FOR UPDATE to acquire a row-level lock
101 stmt = select(EntityLabelsModel).where(EntityLabelsModel.entity_key == entity_key).with_for_update()
102 result = session.scalars(stmt).first()
103
104 if result is None:
105 labels = EntityLabels()
106 else:
107 labels = EntityLabels.deserialize(result.labels)
108
109 # Yield control - The default LabelsProvider will modify the labels IN PLACE
110 yield labels
111
112 # After yield, write the modified labels back
113 labels_dict = labels.serialize()
114 upsert_stmt = insert(EntityLabelsModel).values(entity_key=entity_key, labels=labels_dict)
115 upsert_stmt = upsert_stmt.on_conflict_do_update(
116 index_elements=['entity_key'], set_={EntityLabelsModel.labels: labels_dict}
117 )
118 session.execute(upsert_stmt)
119
120 session.commit()
121 logger.debug(f'Committed atomic read-modify-write for entity {entity_key}', labels_dict)
122
123 except Exception:
124 session.rollback()
125 logger.error(f'Rolled back atomic read-modify-write for entity {entity_key}')
126 raise