Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .github/copilot-instructions.md
Original file line number Diff line number Diff line change
Expand Up @@ -53,7 +53,7 @@ python -m twine upload dist/* # Publish to PyPI

### Version & Dependencies Management
- Version defined in `mixpanel/__init__.py` as `__version__`
- Uses Pydantic v2+ for data validation (`mixpanel/flags/types.py`)
- Feature flag types are stdlib dataclasses with `from_dict` parsers (`mixpanel/flags/types.py`); no third-party validation library
- json-logic library for runtime flag evaluation rules

## Feature Flag Specifics
Expand Down
3 changes: 1 addition & 2 deletions CLAUDE.md
Original file line number Diff line number Diff line change
Expand Up @@ -70,7 +70,7 @@ python -m ghp_import -n -p docs/_build/html
- `LocalFeatureFlagsProvider`: Client-side evaluation with polling (default 60s interval)
- `RemoteFeatureFlagsProvider`: Server-side evaluation via API calls
- Both providers support async operations
- Types defined in `mixpanel/flags/types.py` using Pydantic models
- Types defined in `mixpanel/flags/types.py` as stdlib dataclasses; API payloads are parsed via each type's `from_dict` classmethod

### Key Design Patterns

Expand All @@ -96,7 +96,6 @@ python -m ghp_import -n -p docs/_build/html

- `requests>=2.4.2, <3`: HTTP client (sync)
- `httpx>=0.27.0`: HTTP client (async)
- `pydantic>=2.0.0`: Data validation and types
- `asgiref>=3.0.0`: Async utilities
- `json-logic>=0.7.0a0`: Runtime rules evaluation

Expand Down
6 changes: 3 additions & 3 deletions mixpanel/flags/local_feature_flags.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from __future__ import annotations

import asyncio
import copy
import logging
import threading
import time
Expand Down Expand Up @@ -322,8 +323,7 @@ def _get_assigned_variant(
variant_hash = normalized_hash(str(context_value), salt)

variants = [
variant.model_copy(deep=True)
for variant in flag_definition.ruleset.variants
copy.deepcopy(variant) for variant in flag_definition.ruleset.variants
]
if rollout.variant_splits:
for variant in variants:
Expand Down Expand Up @@ -501,7 +501,7 @@ def _handle_response(
flags = {}
try:
json_data = response.json()
experimentation_flags = ExperimentationFlags.model_validate(json_data)
experimentation_flags = ExperimentationFlags.from_dict(json_data)
for flag in experimentation_flags.flags:
flag.ruleset.variants.sort(key=lambda variant: variant.key)
flags[flag.key] = flag
Expand Down
2 changes: 1 addition & 1 deletion mixpanel/flags/remote_feature_flags.py
Original file line number Diff line number Diff line change
Expand Up @@ -379,7 +379,7 @@ def _build_tracking_properties(

def _handle_response(self, response: httpx.Response) -> dict[str, SelectedVariant]:
response.raise_for_status()
flags_response = RemoteFlagsResponse.model_validate(response.json())
flags_response = RemoteFlagsResponse.from_dict(response.json())
return flags_response.flags

@staticmethod
Expand Down
5 changes: 3 additions & 2 deletions mixpanel/flags/test_local_feature_flags.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import asyncio
import threading
from concurrent.futures import ThreadPoolExecutor
from dataclasses import asdict
from itertools import chain, repeat
from typing import Any
from unittest.mock import Mock, patch
Expand Down Expand Up @@ -85,7 +86,7 @@ def create_test_flag(
def create_flags_response(flags: list[ExperimentationFlag]) -> httpx.Response:
if flags is None:
flags = []
response_data = ExperimentationFlags(flags=flags).model_dump()
response_data = asdict(ExperimentationFlags(flags=flags))
return httpx.Response(status_code=200, json=response_data)


Expand Down Expand Up @@ -809,7 +810,7 @@ async def test_track_exposure_event_successfully_tracks(self):
flag = create_test_flag()
await self.setup_flags([flag])

variant = SelectedVariant(key="treatment", variant_value="treatment")
variant = SelectedVariant(variant_key="treatment", variant_value="treatment")
self._flags.track_exposure_event(TEST_FLAG_KEY, variant, USER_CONTEXT)

self._mock_tracker.assert_called_once()
Expand Down
7 changes: 4 additions & 3 deletions mixpanel/flags/test_remote_feature_flags.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import asyncio
import threading
from concurrent.futures import ThreadPoolExecutor
from dataclasses import asdict
from unittest.mock import Mock

import httpx
Expand All @@ -25,9 +26,9 @@
def create_success_response(
assigned_variants_per_flag: dict[str, SelectedVariant],
) -> httpx.Response:
serialized_response = RemoteFlagsResponse(
code=200, flags=assigned_variants_per_flag
).model_dump()
serialized_response = asdict(
RemoteFlagsResponse(code=200, flags=assigned_variants_per_flag)
)
return httpx.Response(status_code=200, json=serialized_response)


Expand Down
146 changes: 120 additions & 26 deletions mixpanel/flags/types.py
Original file line number Diff line number Diff line change
@@ -1,14 +1,30 @@
from concurrent.futures import Executor
from typing import Any, Literal, Optional

from pydantic import BaseModel, ConfigDict
from dataclasses import dataclass, fields, replace
from typing import Any, Literal, Optional, TypeVar

MIXPANEL_DEFAULT_API_ENDPOINT = "api.mixpanel.com"

_Dataclass = TypeVar("_Dataclass")


def _from_payload(
cls: type[_Dataclass], payload: dict[str, Any], **parsed_fields: Any
) -> _Dataclass:
"""Build a dataclass from an API payload.

Keys the dataclass does not declare are ignored so that new fields added
to the API response never break older SDK versions. Missing required
fields raise TypeError from the dataclass constructor. ``parsed_fields``
override raw payload values for fields that hold nested models.
"""
declared_fields = {field.name for field in fields(cls)} # type: ignore[arg-type]
kwargs = {key: value for key, value in payload.items() if key in declared_fields}
kwargs.update(parsed_fields)
return cls(**kwargs)

class FlagsConfig(BaseModel):
model_config = ConfigDict(arbitrary_types_allowed=True)

@dataclass
class FlagsConfig:
api_host: str = "api.mixpanel.com"
request_timeout_in_seconds: int = 10
# Optional executor used to dispatch exposure-event HTTP sends so flag
Expand All @@ -17,45 +33,91 @@ class FlagsConfig(BaseModel):
exposure_executor: Optional[Executor] = None


@dataclass
class LocalFlagsConfig(FlagsConfig):
enable_polling: bool = True
polling_interval_in_seconds: int = 60


@dataclass
class RemoteFlagsConfig(FlagsConfig):
pass


class Variant(BaseModel):
@dataclass
class Variant:
key: str
value: Any
is_control: bool
split: Optional[float] = 0.0

@classmethod
def from_dict(cls, payload: dict[str, Any]) -> "Variant":
return _from_payload(cls, payload)


class FlagTestUsers(BaseModel):
@dataclass
class FlagTestUsers:
users: dict[str, str]

@classmethod
def from_dict(cls, payload: dict[str, Any]) -> "FlagTestUsers":
return _from_payload(cls, payload)


class VariantOverride(BaseModel):
@dataclass
class VariantOverride:
key: str

@classmethod
def from_dict(cls, payload: dict[str, Any]) -> "VariantOverride":
return _from_payload(cls, payload)


class Rollout(BaseModel):
@dataclass
class Rollout:
rollout_percentage: float
runtime_evaluation_definition: Optional[dict[str, str]] = None
runtime_evaluation_rule: Optional[dict[Any, Any]] = None
variant_override: Optional[VariantOverride] = None
variant_splits: Optional[dict[str, float]] = None

@classmethod
def from_dict(cls, payload: dict[str, Any]) -> "Rollout":
variant_override = payload.get("variant_override")
return _from_payload(
cls,
payload,
variant_override=(
VariantOverride.from_dict(variant_override)
if variant_override is not None
else None
),
)


class RuleSet(BaseModel):
@dataclass
class RuleSet:
variants: list[Variant]
rollout: list[Rollout]
test: Optional[FlagTestUsers] = None

@classmethod
def from_dict(cls, payload: dict[str, Any]) -> "RuleSet":
test_users = payload.get("test")
return _from_payload(
cls,
payload,
variants=[Variant.from_dict(variant) for variant in payload["variants"]],
rollout=[Rollout.from_dict(rollout) for rollout in payload["rollout"]],
test=(
FlagTestUsers.from_dict(test_users) if test_users is not None else None
),
)


class ExperimentationFlag(BaseModel):
@dataclass
class ExperimentationFlag:
id: str
name: str
key: str
Expand All @@ -67,6 +129,12 @@ class ExperimentationFlag(BaseModel):
is_experiment_active: Optional[bool] = None
hash_salt: Optional[str] = None

@classmethod
def from_dict(cls, payload: dict[str, Any]) -> "ExperimentationFlag":
return _from_payload(
cls, payload, ruleset=RuleSet.from_dict(payload["ruleset"])
)


class VariantSource:
"""Where a SelectedVariant came from.
Expand All @@ -81,7 +149,8 @@ class VariantSource:
FALLBACK = "fallback"


class FallbackReason(BaseModel):
@dataclass(frozen=True)
class FallbackReason:
"""Why the SDK returned the developer fallback.

Only meaningful when SelectedVariant.variant_source == VariantSource.FALLBACK.
Expand All @@ -93,8 +162,6 @@ class FallbackReason(BaseModel):
FlagResolutionDetails.error_message.
"""

model_config = ConfigDict(frozen=True)

kind: Literal[
"FLAG_NOT_FOUND",
"MISSING_CONTEXT_KEY",
Expand Down Expand Up @@ -130,40 +197,67 @@ def backend_error(cls, message: str) -> "FallbackReason":
_NO_ROLLOUT_MATCH = FallbackReason(kind="NO_ROLLOUT_MATCH")


class SelectedVariant(BaseModel):
@dataclass
class SelectedVariant:
variant_value: Any
# variant_key can be None if being used as a fallback
variant_key: Optional[str] = None
variant_value: Any
experiment_id: Optional[str] = None
is_experiment_active: Optional[bool] = None
is_qa_tester: Optional[bool] = None
variant_source: Optional[str] = None
# None on success; set when variant_source == FALLBACK
fallback_reason: Optional[FallbackReason] = None

@classmethod
def from_dict(cls, payload: dict[str, Any]) -> "SelectedVariant":
fallback_reason = payload.get("fallback_reason")
return _from_payload(
cls,
payload,
fallback_reason=(
_from_payload(FallbackReason, fallback_reason)
if fallback_reason is not None
else None
),
)

def with_source(self, source: str) -> "SelectedVariant":
"""Return a copy of this variant tagged with the given source.

Clears fallback_reason — use as_fallback() if returning a fallback.
"""
return self.model_copy(
update={"variant_source": source, "fallback_reason": None}
)
return replace(self, variant_source=source, fallback_reason=None)

def as_fallback(self, reason: FallbackReason) -> "SelectedVariant":
"""Return a copy of this variant tagged as a fallback with the given reason."""
return self.model_copy(
update={
"variant_source": VariantSource.FALLBACK,
"fallback_reason": reason,
}
return replace(
self, variant_source=VariantSource.FALLBACK, fallback_reason=reason
)


class ExperimentationFlags(BaseModel):
@dataclass
class ExperimentationFlags:
flags: list[ExperimentationFlag]

@classmethod
def from_dict(cls, payload: dict[str, Any]) -> "ExperimentationFlags":
return cls(
flags=[ExperimentationFlag.from_dict(flag) for flag in payload["flags"]]
)


class RemoteFlagsResponse(BaseModel):
@dataclass
class RemoteFlagsResponse:
code: int
flags: dict[str, SelectedVariant]

@classmethod
def from_dict(cls, payload: dict[str, Any]) -> "RemoteFlagsResponse":
return cls(
code=payload["code"],
flags={
flag_key: SelectedVariant.from_dict(variant)
for flag_key, variant in payload["flags"].items()
},
)
8 changes: 4 additions & 4 deletions openfeature-provider/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,7 @@ from openfeature import api
# 1. Create and register the provider with local evaluation
provider = MixpanelProvider.from_local_config(
"YOUR_PROJECT_TOKEN",
LocalFlagsConfig(token="YOUR_PROJECT_TOKEN"),
LocalFlagsConfig(),
)
api.set_provider(provider)

Expand All @@ -65,7 +65,7 @@ from mixpanel.flags.types import LocalFlagsConfig

provider = MixpanelProvider.from_local_config(
"YOUR_PROJECT_TOKEN",
LocalFlagsConfig(token="YOUR_PROJECT_TOKEN"),
LocalFlagsConfig(),
)
```

Expand All @@ -81,7 +81,7 @@ from mixpanel.flags.types import RemoteFlagsConfig

provider = MixpanelProvider.from_remote_config(
"YOUR_PROJECT_TOKEN",
RemoteFlagsConfig(token="YOUR_PROJECT_TOKEN"),
RemoteFlagsConfig(),
)
```

Expand All @@ -95,7 +95,7 @@ from mixpanel.flags.types import LocalFlagsConfig
from mixpanel_openfeature import MixpanelProvider

# Your existing Mixpanel instance
mp = Mixpanel("YOUR_PROJECT_TOKEN", local_flags_config=LocalFlagsConfig(token="YOUR_PROJECT_TOKEN"))
mp = Mixpanel("YOUR_PROJECT_TOKEN", local_flags_config=LocalFlagsConfig())
local_flags = mp.local_flags
local_flags.start_polling_for_definitions()

Expand Down
Loading
Loading