Mirrored from GitHub github.com/roostorg/osprey
0

Configure Feed

Select the types of activity you want to include in your feed.

osprey / example_plugins / src / services / labels_service.py
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