Skip to content

Commit 586e4d2

Browse files
ChromeHeartsOrbax Authors
authored andcommitted
Internal
PiperOrigin-RevId: 925535181
1 parent 8ac9c61 commit 586e4d2

12 files changed

Lines changed: 2500 additions & 10 deletions

File tree

checkpoint/orbax/checkpoint/experimental/tiering_service/db_lib.py

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,8 +15,13 @@
1515
"""Database initialization utilities for Tiering Service."""
1616

1717
import contextlib
18+
import sqlite3
19+
1820
from orbax.checkpoint.experimental.tiering_service import db_schema
1921
from orbax.checkpoint.experimental.tiering_service.proto import tiering_service_pb2
22+
from sqlalchemy import event
23+
from sqlalchemy.dialects.sqlite.aiosqlite import AsyncAdapt_aiosqlite_connection
24+
from sqlalchemy.engine import Engine
2025
from sqlalchemy.exc import OperationalError
2126
from sqlalchemy.ext.asyncio import AsyncEngine
2227
from sqlalchemy.ext.asyncio import AsyncSession
@@ -25,6 +30,29 @@
2530
from sqlalchemy.orm import sessionmaker
2631

2732

33+
@event.listens_for(Engine, "connect")
34+
def set_sqlite_pragma(dbapi_connection, connection_record):
35+
"""Enables foreign key constraints on SQLite database connections.
36+
37+
This is SQLite-specific because other databases (like PostgreSQL) enforce
38+
foreign keys by default and do not support SQLite's PRAGMA syntax.
39+
We perform an isinstance check against the standard sqlite3.Connection
40+
and SQLAlchemy's aiosqlite adapter wrapper to verify if this is an SQLite
41+
connection.
42+
43+
Args:
44+
dbapi_connection: The database connection to configure.
45+
connection_record: Metadata about the connection.
46+
"""
47+
del connection_record
48+
connection_types = (sqlite3.Connection, AsyncAdapt_aiosqlite_connection)
49+
50+
if isinstance(dbapi_connection, connection_types):
51+
cursor = dbapi_connection.cursor()
52+
cursor.execute("PRAGMA foreign_keys=ON")
53+
cursor.close()
54+
55+
2856
def get_async_engine(config: tiering_service_pb2.ServerConfig) -> AsyncEngine:
2957
"""Returns an AsyncEngine configured from ServerConfig."""
3058
input_url = config.db_connection_str

checkpoint/orbax/checkpoint/experimental/tiering_service/db_schema.py

Lines changed: 61 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -63,6 +63,17 @@ class RequestType(enum.IntEnum):
6363
REQUEST_TYPE_DELETE_FROM_ALL_TIERS = 3
6464

6565

66+
class TierPathState(enum.IntEnum):
67+
"""The state of an asset's storage location (tier path)."""
68+
69+
UNSPECIFIED = 0
70+
PENDING = 1
71+
IN_PROGRESS = 2
72+
READY = 3
73+
FAILED = 4
74+
DELETED = 5
75+
76+
6677
class Asset(Base):
6778
"""A CTS asset representing a complete checkpoint.
6879
@@ -331,26 +342,44 @@ class TierPath(Base):
331342
nullable=False,
332343
default=lambda: str(uuid.uuid4()),
333344
)
345+
state = sqlalchemy.Column(
346+
sqlalchemy.Enum(TierPathState),
347+
default=TierPathState.PENDING,
348+
nullable=False,
349+
)
334350

335351
asset = sqlalchemy.orm.relationship("Asset", back_populates="tier_paths")
336352
storage_backend = sqlalchemy.orm.relationship(
337353
"StorageBackend", back_populates="tier_paths"
338354
)
339355

340356
__table_args__ = (
341-
# An asset can have at most one TierPath for a given storage backend.
342-
sqlalchemy.UniqueConstraint(
357+
# Enforce uniqueness of (asset, backend) only for active tier paths
358+
# (PENDING, IN_PROGRESS, READY).
359+
sqlalchemy.Index(
360+
"uq_tier_path_active_backend",
343361
"asset_uuid",
344362
"storage_backend_id",
345-
name="uq_tier_path_asset_backend",
363+
unique=True,
364+
sqlite_where=sqlalchemy.column("state").in_([
365+
TierPathState.PENDING.name,
366+
TierPathState.IN_PROGRESS.name,
367+
TierPathState.READY.name,
368+
]),
369+
postgresql_where=sqlalchemy.column("state").in_([
370+
TierPathState.PENDING.name,
371+
TierPathState.IN_PROGRESS.name,
372+
TierPathState.READY.name,
373+
]),
346374
),
347375
)
348376

349377
def __repr__(self):
350378
return (
351379
f"TierPath(id={self.id}, asset_uuid='{self.asset_uuid}',"
352380
f" storage_backend_id={self.storage_backend_id}, path='{self.path}',"
353-
f" ready_at={self.ready_at}, expires_at={self.expires_at})"
381+
f" state={self.state.name}, ready_at={self.ready_at},"
382+
f" expires_at={self.expires_at})"
354383
)
355384

356385

@@ -367,6 +396,13 @@ class AssetJob(Base):
367396
status: Current execution status of the job, an instance of JobStatus.
368397
target_tier_path_id: Foreign key to the targeted TierPath for operations
369398
such as COPY or DELETE_FROM_INSTANCE.
399+
request_id: A unique identifier (UUID) for this job execution request.
400+
transfer_status: JSON dictionary containing progress and GCP operation
401+
details.
402+
expiration_at: Timestamp when the worker's lease on this job expires.
403+
last_updated_at: Timestamp of the last status update or heartbeat.
404+
worker_host: Hostname of the worker processing this job.
405+
worker_pid: Process ID of the worker processing this job.
370406
created_at: Timestamp when the job was created.
371407
completed_at: Timestamp when the job was completed.
372408
asset: Relationship to the associated Asset.
@@ -396,9 +432,24 @@ class AssetJob(Base):
396432
# Target tier path for COPY and DELETE_FROM_INSTANCE requests
397433
target_tier_path_id = sqlalchemy.Column(
398434
sqlalchemy.Integer,
399-
sqlalchemy.ForeignKey("tier_paths.id", ondelete="CASCADE"),
435+
sqlalchemy.ForeignKey("tier_paths.id"),
400436
nullable=True,
401437
)
438+
request_id = sqlalchemy.Column(
439+
sqlalchemy.String,
440+
nullable=False,
441+
unique=True,
442+
default=lambda: str(uuid.uuid4()),
443+
)
444+
transfer_status = sqlalchemy.Column(sqlalchemy.JSON, nullable=True)
445+
expiration_at = sqlalchemy.Column(
446+
sqlalchemy.DateTime(timezone=True), nullable=True
447+
)
448+
last_updated_at = sqlalchemy.Column(
449+
sqlalchemy.DateTime(timezone=True), nullable=True
450+
)
451+
worker_host = sqlalchemy.Column(sqlalchemy.String, nullable=True)
452+
worker_pid = sqlalchemy.Column(sqlalchemy.Integer, nullable=True)
402453

403454
created_at = sqlalchemy.Column(
404455
sqlalchemy.DateTime(timezone=True),
@@ -416,9 +467,11 @@ class AssetJob(Base):
416467
# target_tier_path is required in COPY and DELETE_FROM_INSTANCE requests.
417468
sqlalchemy.CheckConstraint(
418469
"""
419-
(request_type IN ('REQUEST_TYPE_COPY', 'REQUEST_TYPE_DELETE_FROM_INSTANCE') AND target_tier_path_id IS NOT NULL)
470+
(request_type IN ('REQUEST_TYPE_COPY', 'REQUEST_TYPE_DELETE_FROM_INSTANCE')
471+
AND target_tier_path_id IS NOT NULL)
420472
OR
421-
(request_type IN ('REQUEST_TYPE_DELETE_FROM_ALL_TIERS', 'REQUEST_TYPE_UNSPECIFIED') AND target_tier_path_id IS NULL)
473+
(request_type IN ('REQUEST_TYPE_DELETE_FROM_ALL_TIERS', 'REQUEST_TYPE_UNSPECIFIED')
474+
AND target_tier_path_id IS NULL)
422475
""",
423476
name="check_asset_job_valid_payload",
424477
),
@@ -430,5 +483,6 @@ def __repr__(self):
430483
f" request_type='{self.request_type.name}',"
431484
f" status='{self.status.name}',"
432485
f" target_tier_path_id={self.target_tier_path_id},"
486+
f" request_id='{self.request_id}',"
433487
f" created_at={self.created_at}, completed_at={self.completed_at})"
434488
)

checkpoint/orbax/checkpoint/experimental/tiering_service/db_schema_test.py

Lines changed: 119 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,12 +22,22 @@
2222
import greenlet # pylint: disable=unused-import
2323
from orbax.checkpoint.experimental.tiering_service import db_schema
2424
import sqlalchemy
25+
from sqlalchemy import event
26+
from sqlalchemy.engine import Engine
2527
from sqlalchemy.ext.asyncio import AsyncSession
2628
from sqlalchemy.ext.asyncio import create_async_engine
2729
from sqlalchemy.future import select
2830
from sqlalchemy.orm import sessionmaker
2931

3032

33+
@event.listens_for(Engine, "connect")
34+
def set_sqlite_pragma(dbapi_connection, connection_record):
35+
del connection_record
36+
cursor = dbapi_connection.cursor()
37+
cursor.execute("PRAGMA foreign_keys=ON")
38+
cursor.close()
39+
40+
3141
class DbSchemaTest(parameterized.TestCase, unittest.IsolatedAsyncioTestCase):
3242

3343
async def asyncSetUp(self) -> None:
@@ -442,6 +452,115 @@ async def test_asset_job_queue(self) -> None:
442452
fetched_job.status, db_schema.JobStatus.JOB_STATUS_COMPLETED
443453
)
444454

455+
async def test_tier_path_deletion_fails_when_referenced_by_job(self) -> None:
456+
async with self.session_maker() as session:
457+
asset = db_schema.Asset(
458+
asset_uuid="uuid-delete-tp-fail",
459+
path="/experiment/delete-tp-fail",
460+
user="testuser",
461+
)
462+
backend = db_schema.StorageBackend(
463+
level=0,
464+
zone="us-central1-a",
465+
backend_type=db_schema.BackendType.BACKEND_TYPE_GCS,
466+
prefix="gs://gcs-bucket",
467+
)
468+
tier_path = db_schema.TierPath(
469+
asset_uuid="uuid-delete-tp-fail",
470+
storage_backend=backend,
471+
path="/path1",
472+
)
473+
session.add_all([asset, backend, tier_path])
474+
await session.commit()
475+
476+
job = db_schema.AssetJob(
477+
asset_uuid="uuid-delete-tp-fail",
478+
request_type=db_schema.RequestType.REQUEST_TYPE_COPY,
479+
status=db_schema.JobStatus.JOB_STATUS_COMPLETED,
480+
target_tier_path_id=tier_path.id,
481+
)
482+
session.add(job)
483+
await session.commit()
484+
485+
# Deleting the tier path should fail with IntegrityError because the job
486+
# still references it.
487+
await session.delete(tier_path)
488+
with self.assertRaises(sqlalchemy.exc.IntegrityError):
489+
await session.commit()
490+
491+
async def test_tier_path_conditional_uniqueness(self) -> None:
492+
async with self.session_maker() as session:
493+
asset = db_schema.Asset(
494+
asset_uuid="uuid-cond-uniq",
495+
path="/experiment/cond-uniq",
496+
user="testuser",
497+
)
498+
backend = db_schema.StorageBackend(
499+
level=0,
500+
zone="us-central1-a",
501+
backend_type=db_schema.BackendType.BACKEND_TYPE_GCS,
502+
prefix="gs://gcs-bucket",
503+
)
504+
session.add_all([asset, backend])
505+
await session.commit()
506+
507+
# 1. We can insert one PENDING tier path
508+
tp_pending = db_schema.TierPath(
509+
asset_uuid="uuid-cond-uniq",
510+
storage_backend=backend,
511+
path="/path-pending",
512+
state=db_schema.TierPathState.PENDING,
513+
)
514+
session.add(tp_pending)
515+
await session.commit()
516+
517+
# 2. Trying to insert another IN_PROGRESS tier path should fail
518+
tp_in_progress = db_schema.TierPath(
519+
asset_uuid="uuid-cond-uniq",
520+
storage_backend=backend,
521+
path="/path-inprogress",
522+
state=db_schema.TierPathState.IN_PROGRESS,
523+
)
524+
session.add(tp_in_progress)
525+
with self.assertRaises(sqlalchemy.exc.IntegrityError):
526+
await session.commit()
527+
await session.rollback()
528+
529+
# 3. Transition the first one to FAILED
530+
result = await session.execute(
531+
select(db_schema.TierPath).filter_by(asset_uuid="uuid-cond-uniq")
532+
)
533+
tp = result.scalars().first()
534+
tp.state = db_schema.TierPathState.FAILED
535+
await session.commit()
536+
537+
# 4. Now we can insert a new PENDING tier path for the same backend!
538+
tp_new_pending = db_schema.TierPath(
539+
asset_uuid="uuid-cond-uniq",
540+
storage_backend=backend,
541+
path="/path-new-pending",
542+
state=db_schema.TierPathState.PENDING,
543+
)
544+
session.add(tp_new_pending)
545+
await session.commit()
546+
547+
# 5. And we can also insert a DELETED tier path for the same backend!
548+
tp_deleted = db_schema.TierPath(
549+
asset_uuid="uuid-cond-uniq",
550+
storage_backend=backend,
551+
path="/path-deleted",
552+
state=db_schema.TierPathState.DELETED,
553+
)
554+
session.add(tp_deleted)
555+
await session.commit()
556+
557+
# Verify all 3 rows exist (FAILED, DELETED, PENDING)
558+
result = await session.execute(
559+
select(db_schema.TierPath).filter_by(asset_uuid="uuid-cond-uniq")
560+
)
561+
paths = result.scalars().all()
562+
self.assertLen(paths, 3)
563+
445564
async def test_create_asset_duplicates_allowed_for_deleted_incomplete(self):
446565
# Verify we can have duplicate path for DELETED or INCOMPLETE states
447566
async with self.session_maker() as session:

0 commit comments

Comments
 (0)