Coverage for polar/kit/trial.py: 67%
40 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 12:42 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 12:42 +0000
1from datetime import datetime
2from enum import StrEnum
3from typing import Self
5from dateutil.relativedelta import relativedelta
6from pydantic import BaseModel, Field, model_validator
7from pydantic_core import PydanticCustomError
8from sqlalchemy import Integer
9from sqlalchemy.orm import Mapped, mapped_column
11from polar.kit.extensions.sqlalchemy.types import StringEnum
14class TrialInterval(StrEnum):
15 day = "day"
16 week = "week"
17 month = "month"
18 year = "year"
20 def get_end(self, d: datetime, count: int) -> datetime:
21 match self:
22 case TrialInterval.day:
23 return d + relativedelta(days=count)
24 case TrialInterval.week:
25 return d + relativedelta(weeks=count)
26 case TrialInterval.month:
27 return d + relativedelta(months=count)
28 case TrialInterval.year:
29 return d + relativedelta(years=count)
32class TrialConfigurationMixin:
33 trial_interval: Mapped[TrialInterval | None] = mapped_column(
34 StringEnum(TrialInterval), nullable=True, default=None
35 )
36 trial_interval_count: Mapped[int | None] = mapped_column(
37 Integer, nullable=True, default=None
38 )
41class TrialConfigurationInputMixin(BaseModel):
42 trial_interval: TrialInterval | None = Field(
43 default=None, description="The interval unit for the trial period."
44 )
45 trial_interval_count: int | None = Field(
46 default=None,
47 description="The number of interval units for the trial period.",
48 ge=1,
49 le=1000,
50 )
52 @model_validator(mode="after")
53 def is_complete_configuration(self) -> Self:
54 if self.trial_interval is None and self.trial_interval_count is None:
55 return self
57 if self.trial_interval is not None and self.trial_interval_count is not None:
58 return self
60 raise PydanticCustomError(
61 "missing",
62 "Both trial_interval and trial_interval_count must be set together.",
63 )
66class TrialConfigurationOutputMixin(BaseModel):
67 trial_interval: TrialInterval | None = Field(
68 description="The interval unit for the trial period."
69 )
70 trial_interval_count: int | None = Field(
71 description="The number of interval units for the trial period."
72 )