aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--generalresearch/managers/thl/user_manager/user_metadata_manager.py23
-rw-r--r--generalresearch/models/custom_types.py28
-rw-r--r--tests/managers/thl/test_user_manager/test_user_metadata.py12
3 files changed, 57 insertions, 6 deletions
diff --git a/generalresearch/managers/thl/user_manager/user_metadata_manager.py b/generalresearch/managers/thl/user_manager/user_metadata_manager.py
index 2b18302..1d82a2b 100644
--- a/generalresearch/managers/thl/user_manager/user_metadata_manager.py
+++ b/generalresearch/managers/thl/user_manager/user_metadata_manager.py
@@ -2,7 +2,11 @@ from __future__ import annotations
from collections.abc import Collection
+import email_normalize
+from pydantic import EmailStr, TypeAdapter
+
from generalresearch.managers.base import PostgresManager
+from generalresearch.models.custom_types import CanonicalEmailStr
from generalresearch.models.thl.user_profile import UserMetadata
@@ -26,6 +30,20 @@ class UserMetadataManager(PostgresManager):
return {x["product_user_id"]: UserMetadata.from_db(**x) for x in res}
+ def filter_by_email_aliases(
+ self, email_addresses: Collection[EmailStr]
+ ) -> list[UserMetadata]:
+ """
+ Accepts raw or canonical emails, normalizes them, and then filters
+ by the normalized/canonical email.
+ """
+ validated_emails = TypeAdapter(list[EmailStr]).validate_python(email_addresses)
+ canonical_emails = {
+ email_normalize.normalize(email).normalized_address
+ for email in validated_emails
+ }
+ return self.filter(canonical_emails=canonical_emails)
+
def filter(
self,
user_ids: Collection[int] | None = None,
@@ -33,7 +51,7 @@ class UserMetadataManager(PostgresManager):
email_sha256s: Collection[str] | None = None,
email_sha1s: Collection[str] | None = None,
email_md5s: Collection[str] | None = None,
- canonical_emails: Collection[str] | None = None,
+ canonical_emails: Collection[CanonicalEmailStr] | None = None,
) -> list[UserMetadata]:
for arg in [
user_ids,
@@ -66,6 +84,9 @@ class UserMetadataManager(PostgresManager):
params["email_md5"] = list(set(email_md5s))
filters.append("email_md5 = ANY(%(email_md5)s)")
if canonical_emails is not None:
+ canonical_emails = TypeAdapter(list[CanonicalEmailStr]).validate_python(
+ canonical_emails
+ )
params["canonical_emails"] = list(set(canonical_emails))
filters.append("canonical_email = ANY(%(canonical_emails)s)")
diff --git a/generalresearch/models/custom_types.py b/generalresearch/models/custom_types.py
index 680a99c..0040b92 100644
--- a/generalresearch/models/custom_types.py
+++ b/generalresearch/models/custom_types.py
@@ -6,9 +6,11 @@ from datetime import UTC, datetime, timedelta
from typing import Annotated, Any, Literal
from uuid import UUID
+import email_normalize
from pydantic import (
AnyUrl,
AwareDatetime,
+ EmailStr,
Field,
HttpUrl,
IPvAnyAddress,
@@ -36,6 +38,22 @@ def validate_hostname(v: str) -> str:
InternalHostname = Annotated[str, AfterValidator(validate_hostname)]
+def validate_canonical_email(value: str) -> str:
+ normalized_email = email_normalize.normalize(value).normalized_address
+ if value != normalized_email:
+ raise ValueError(
+ f"canonical email must already be normalized: "
+ f"{value!r} != {normalized_email!r}"
+ )
+ return value
+
+
+CanonicalEmailStr = Annotated[
+ EmailStr,
+ BeforeValidator(validate_canonical_email),
+]
+
+
class PostgresDict(MultiHostHost):
"""The path part of this host, or `None`."""
@@ -62,9 +80,9 @@ def convert_str_dt(v: Any) -> AwareDatetime | None:
def assert_utc(v: AwareDatetime) -> AwareDatetime:
if isinstance(v, datetime):
# We need utcoffset b/c FastAPI parses datetimes using FixedTimezone
- assert v.tzinfo == UTC or v.tzinfo.utcoffset(v) == timedelta(
- 0
- ), "Timezone is not UTC"
+ assert v.tzinfo == UTC or v.tzinfo.utcoffset(v) == timedelta(0), (
+ "Timezone is not UTC"
+ )
v = v.astimezone(UTC)
return v
@@ -98,7 +116,7 @@ LanguageISOLike = Annotated[
def check_valid_uuid(v: str) -> str:
try:
assert UUID(v).hex == v
- except (ValueError, AssertionError):
+ except ValueError, AssertionError:
raise ValueError("Invalid UUID")
return v
@@ -106,7 +124,7 @@ def check_valid_uuid(v: str) -> str:
def is_valid_uuid(v: str) -> bool:
try:
assert UUID(v).hex == v
- except (ValueError, AssertionError):
+ except ValueError, AssertionError:
return False
return True
diff --git a/tests/managers/thl/test_user_manager/test_user_metadata.py b/tests/managers/thl/test_user_manager/test_user_metadata.py
index 515c5bc..7e745b6 100644
--- a/tests/managers/thl/test_user_manager/test_user_metadata.py
+++ b/tests/managers/thl/test_user_manager/test_user_metadata.py
@@ -141,3 +141,15 @@ class TestUserMetadataManager:
res = user_metadata_manager.filter(canonical_emails=[expected_canonical])
assert len(res) == 1
+
+ res = user_metadata_manager.filter_by_email_aliases(
+ email_addresses=[f"{local}+789@googlemail.com", expected_canonical]
+ )
+ assert len(res) == 1
+
+ with pytest.raises(
+ ValueError, match="canonical email must already be normalized"
+ ):
+ user_metadata_manager.filter(
+ canonical_emails=[f"{local}+789@googlemail.com"]
+ )