Files
agentrhq__authsome/tests/auth/test_models.py
2026-06-04 22:22:50 +05:30

332 lines
11 KiB
Python

"""Tests for authsome data models."""
# ruff: noqa: PLR2004
from authsome.auth.models.config import current_spec_version
from authsome.auth.models.connection import (
ConnectionRecord,
ProviderClientRecord,
ProviderMetadataRecord,
ProviderStateRecord,
Sensitive,
)
from authsome.auth.models.enums import AuthType, ConnectionStatus, ExportFormat, FlowType
from authsome.auth.models.provider import ApiKeyConfig, OAuthConfig, ProviderDefinition
from authsome.errors import OperationNotAllowedError
from authsome.identity.helpers import IdentityMetadata
class TestEnums:
"""Enum serialization and values."""
def test_auth_type_values(self) -> None:
assert AuthType.OAUTH2.value == "oauth2"
assert AuthType.API_KEY.value == "api_key"
def test_flow_type_values(self) -> None:
assert FlowType.DCR_PKCE.value == "dcr_pkce"
assert FlowType.API_KEY.value == "api_key"
def test_connection_status_values(self) -> None:
assert ConnectionStatus.CONNECTED.value == "connected"
assert ConnectionStatus.EXPIRED.value == "expired"
assert ConnectionStatus.REVOKED.value == "revoked"
def test_export_format_values(self) -> None:
assert ExportFormat.ENV.value == "env"
assert ExportFormat.JSON.value == "json"
class TestSpecVersion:
"""Spec version helper tests."""
def test_current_spec_version_is_int(self) -> None:
assert isinstance(current_spec_version(), int)
class TestIdentityMetadata:
"""Identity metadata model tests."""
def test_required_fields(self) -> None:
meta = IdentityMetadata(handle="steady-wisely-boldly-0042", did="did:key:z6MkTest")
assert meta.handle == "steady-wisely-boldly-0042"
assert meta.did == "did:key:z6MkTest"
assert meta.created_at is not None
assert meta.updated_at is not None
def test_json_roundtrip(self) -> None:
meta = IdentityMetadata(handle="steady-wisely-boldly-0042", did="did:key:z6MkTest")
json_str = meta.model_dump_json()
restored = IdentityMetadata.model_validate_json(json_str)
assert restored.handle == "steady-wisely-boldly-0042"
assert restored.did == "did:key:z6MkTest"
class TestProviderDefinition:
"""Provider definition model tests."""
def test_oauth_provider(self) -> None:
provider = ProviderDefinition(
name="github",
display_name="GitHub",
auth_type=AuthType.OAUTH2,
flow=FlowType.DCR_PKCE,
oauth=OAuthConfig(
authorization_url="https://github.com/login/oauth/authorize",
token_url="https://github.com/login/oauth/access_token",
scopes=["repo", "read:user"],
),
)
assert provider.auth_type == AuthType.OAUTH2
assert provider.flow == FlowType.DCR_PKCE
assert provider.oauth is not None
assert "repo" in provider.oauth.scopes
def test_api_key_provider(self) -> None:
provider = ProviderDefinition(
name="openai",
display_name="OpenAI",
auth_type=AuthType.API_KEY,
flow=FlowType.API_KEY,
api_key=ApiKeyConfig(
header_name="Authorization",
header_prefix="Bearer",
),
)
assert provider.auth_type == AuthType.API_KEY
assert provider.api_key is not None
def test_json_roundtrip(self) -> None:
provider = ProviderDefinition(
name="test",
display_name="Test",
logo="https://example.com/logo.svg",
description="Example provider for model tests.",
auth_type=AuthType.API_KEY,
flow=FlowType.API_KEY,
api_key=ApiKeyConfig(),
)
json_str = provider.model_dump_json()
restored = ProviderDefinition.model_validate_json(json_str)
assert restored.name == "test"
assert restored.logo == "https://example.com/logo.svg"
assert restored.description == "Example provider for model tests."
def test_docs_field_is_optional_and_parsed(self) -> None:
provider = ProviderDefinition.model_validate(
{
"schema_version": 1,
"name": "calendly",
"display_name": "Calendly",
"auth_type": "api_key",
"flow": "api_key",
"api_key": {"header_name": "Authorization", "header_prefix": "Bearer"},
"docs": "https://example.com/setup",
}
)
assert provider.docs == "https://example.com/setup"
class TestSensitiveAnnotation:
"""Sensitive field annotation tests."""
def test_sensitive_marker_exists(self) -> None:
assert Sensitive is not None
def test_connection_record_has_sensitive_fields(self) -> None:
from authsome.utils import redact
record = ConnectionRecord(
schema_version=2,
provider="openai",
identity="default",
connection_name="default",
auth_type=AuthType.API_KEY,
status=ConnectionStatus.CONNECTED,
api_key="sk-super-secret",
)
data = redact(record)
assert data["api_key"] == "***REDACTED***"
def test_non_sensitive_fields_not_redacted(self) -> None:
from authsome.utils import redact
record = ConnectionRecord(
schema_version=2,
provider="openai",
identity="default",
connection_name="default",
auth_type=AuthType.API_KEY,
status=ConnectionStatus.CONNECTED,
)
data = redact(record)
assert data["provider"] == "openai"
assert data["identity"] == "default"
def test_none_sensitive_fields_not_redacted_when_none(self) -> None:
from authsome.utils import redact
record = ConnectionRecord(
schema_version=2,
provider="openai",
identity="default",
connection_name="default",
auth_type=AuthType.API_KEY,
status=ConnectionStatus.CONNECTED,
api_key=None,
)
data = redact(record)
assert data["api_key"] is None # None stays None, not "***REDACTED***"
def test_provider_client_record_sensitive_fields(self) -> None:
from authsome.utils import redact
record = ProviderClientRecord(
provider="github",
client_id="public-cid",
client_secret="secret-csec",
)
data = redact(record)
assert data["client_secret"] == "***REDACTED***"
assert data["client_id"] == "public-cid"
def test_provider_client_record_is_server_owned(self) -> None:
record = ProviderClientRecord(
provider="github",
client_id="public-cid",
client_secret="secret-csec",
)
data = record.model_dump(mode="json")
assert "identity" not in data
assert data["provider"] == "github"
class TestErrors:
"""Typed error model tests."""
def test_operation_not_allowed_error_includes_context(self) -> None:
error = OperationNotAllowedError(
"revoke",
"Only administrators can revoke providers.",
provider="github",
)
assert str(error) == ("OperationNotAllowedError: [github] (revoke) Only administrators can revoke providers.")
class TestConnectionRecord:
"""Connection record model tests."""
def test_oauth_record(self) -> None:
record = ConnectionRecord(
schema_version=2,
provider="github",
identity="default",
connection_name="personal",
auth_type=AuthType.OAUTH2,
status=ConnectionStatus.CONNECTED,
scopes=["repo"],
token_type="Bearer",
)
assert record.provider == "github"
assert record.status == ConnectionStatus.CONNECTED
def test_api_key_record(self) -> None:
record = ConnectionRecord(
schema_version=2,
provider="openai",
identity="default",
connection_name="default",
auth_type=AuthType.API_KEY,
status=ConnectionStatus.CONNECTED,
)
assert record.auth_type == AuthType.API_KEY
def test_json_roundtrip_with_plaintext_token(self) -> None:
record = ConnectionRecord(
schema_version=2,
provider="test",
identity="default",
connection_name="default",
auth_type=AuthType.OAUTH2,
status=ConnectionStatus.CONNECTED,
access_token="plaintext-token-value",
)
json_str = record.model_dump_json()
restored = ConnectionRecord.model_validate_json(json_str)
assert restored.access_token == "plaintext-token-value"
def test_schema_version_defaults_to_2(self) -> None:
record = ConnectionRecord(
provider="test",
identity="default",
connection_name="default",
auth_type=AuthType.API_KEY,
status=ConnectionStatus.CONNECTED,
)
assert record.schema_version == 2
class TestProviderMetadataRecord:
"""Provider metadata record tests."""
def test_defaults(self) -> None:
meta = ProviderMetadataRecord(identity="default", provider="github")
assert meta.default_connection == "default"
assert meta.connection_names == []
def test_connection_tracking(self) -> None:
meta = ProviderMetadataRecord(
identity="default",
provider="github",
connection_names=["personal", "work"],
last_used_connection="work",
)
assert len(meta.connection_names) == 2
assert meta.last_used_connection == "work"
class TestProviderStateRecord:
"""Provider state record tests."""
def test_defaults(self) -> None:
state = ProviderStateRecord(provider="github", identity="default")
assert state.last_refresh_at is None
assert state.last_refresh_error is None
class TestHostUrl:
"""ProviderDefinition api_url field tests."""
def test_provider_definition_parses_api_url(self) -> None:
provider = ProviderDefinition.model_validate(
{
"schema_version": 1,
"name": "openai",
"display_name": "OpenAI",
"auth_type": "api_key",
"flow": "api_key",
"api_key": {"header_name": "Authorization", "header_prefix": "Bearer"},
"api_url": "api.openai.com",
}
)
assert provider.api_url == "api.openai.com"
def test_provider_definition_defaults_api_url_to_none(self) -> None:
provider = ProviderDefinition.model_validate(
{
"schema_version": 1,
"name": "example",
"display_name": "Example",
"auth_type": "api_key",
"flow": "api_key",
"api_key": {"header_name": "Authorization", "header_prefix": "Bearer"},
}
)
assert provider.api_url is None