This commit is contained in:
Nirvana
2026-01-16 12:24:41 +01:00
parent 34203d2bca
commit eef01919d7
77 changed files with 1910 additions and 2521 deletions
+4 -13
View File
@@ -103,19 +103,12 @@ def get_configured_manager(country: str = "de") -> "ProviderManager":
settings_manager = SettingsManager(enable_kodi_integration=True)
# Detect providers from Kodi (if available) or auto-detect from AVAILABLE_PROVIDERS
if (
settings_manager.kodi_bridge
and settings_manager.kodi_bridge.is_kodi_environment()
):
detected_providers = (
settings_manager.kodi_bridge.detect_all_providers_from_kodi()
)
if settings_manager.kodi_bridge and settings_manager.kodi_bridge.is_kodi_environment():
detected_providers = settings_manager.kodi_bridge.detect_all_providers_from_kodi()
logger.info(f"Detected providers from Kodi: {detected_providers}")
else:
# Auto-generate detected_providers from AVAILABLE_PROVIDERS
logger.info(
f"Not in Kodi environment, auto-detecting providers from AVAILABLE_PROVIDERS"
)
logger.info(f"Not in Kodi environment, auto-detecting providers from AVAILABLE_PROVIDERS")
detected_providers = {}
for provider_name, provider_class in AVAILABLE_PROVIDERS.items():
@@ -132,9 +125,7 @@ def get_configured_manager(country: str = "de") -> "ProviderManager":
logger.debug(f"Provider '{provider_name}' is single-country")
# Use the new discover_providers method with detected providers
registered = manager.discover_providers(
country=country, detected_providers=detected_providers
)
registered = manager.discover_providers(country=country, detected_providers=detected_providers)
logger.info(f"Registered {len(registered)} providers: {registered}")
@@ -1,8 +1,7 @@
# streaming_providers/base/auth/__init__.py
from .base_auth import BaseAuthenticator, BaseAuthToken
from .credential_manager import CredentialManager
from .credentials import (BaseCredentials, ClientCredentials,
UserPasswordCredentials)
from .credentials import BaseCredentials, ClientCredentials, UserPasswordCredentials
from .session_manager import SessionManager
# Only export what consumers should use
+39 -113
View File
@@ -116,13 +116,9 @@ class BaseAuthenticator(ABC):
# Log country configuration
if self.country:
logger.info(
f"Initializing {provider_name} authenticator for country: {country}"
)
logger.info(f"Initializing {provider_name} authenticator for country: {country}")
else:
logger.info(
f"Initializing {provider_name} authenticator (no country specified)"
)
logger.info(f"Initializing {provider_name} authenticator (no country specified)")
# Use injected settings manager or create one for backward compatibility
if settings_manager is not None:
@@ -133,9 +129,7 @@ class BaseAuthenticator(ABC):
self.settings_manager = self._create_settings_manager(
config_dir, enable_kodi_integration
)
logger.info(
f"Created settings manager for backward compatibility for {provider_name}"
)
logger.info(f"Created settings manager for backward compatibility for {provider_name}")
# Register provider with settings manager
if hasattr(self.settings_manager, "register_provider"):
@@ -177,8 +171,7 @@ class BaseAuthenticator(ABC):
):
"""Create fallback using adapter pattern"""
try:
from ..settings.settings_manager_adapter import \
SettingsManagerFactory
from ..settings.settings_manager_adapter import SettingsManagerFactory
adapter = SettingsManagerFactory.create_default_adapter(
prefer_unified=True,
@@ -211,14 +204,10 @@ class BaseAuthenticator(ABC):
def get_provider_credentials(self, provider_name, country=None):
if self.credential_manager:
return self.credential_manager.load_credentials(
provider_name, country
)
return self.credential_manager.load_credentials(provider_name, country)
return None
def save_provider_credentials(
self, provider_name, credentials, country=None
):
def save_provider_credentials(self, provider_name, credentials, country=None):
if self.credential_manager:
return self.credential_manager.save_credentials(
provider_name, credentials, country
@@ -229,9 +218,7 @@ class BaseAuthenticator(ABC):
return self.session_manager.load_token_data(provider_name, country)
def save_token_data(self, provider_name, token_data, country=None):
return self.session_manager.save_session(
provider_name, token_data, country
)
return self.session_manager.save_session(provider_name, token_data, country)
def get_device_id(self, provider_name, country=None):
return self.session_manager.get_device_id(provider_name, country)
@@ -264,9 +251,7 @@ class BaseAuthenticator(ABC):
def get_provider_credentials(self, provider_name, country=None):
return None
def save_provider_credentials(
self, provider_name, credentials, country=None
):
def save_provider_credentials(self, provider_name, credentials, country=None):
return False
def load_token_data(self, provider_name, country=None):
@@ -311,9 +296,7 @@ class BaseAuthenticator(ABC):
logger.debug(
f"Settings manager doesn't support country parameter, using without"
)
return self.settings_manager.get_provider_credentials(
self.provider_name
)
return self.settings_manager.get_provider_credentials(self.provider_name)
else:
logger.warning(
f"Settings manager has no credential loading method for {self.provider_name}"
@@ -321,9 +304,7 @@ class BaseAuthenticator(ABC):
return None
except Exception as e:
logger.error(
f"Error loading credentials from manager for {self.provider_name}: {e}"
)
logger.error(f"Error loading credentials from manager for {self.provider_name}: {e}")
return None
# Abstract methods remain unchanged
@@ -344,9 +325,7 @@ class BaseAuthenticator(ABC):
pass
@abstractmethod
def _create_token_from_response(
self, response_data: Dict[str, Any]
) -> BaseAuthToken:
def _create_token_from_response(self, response_data: Dict[str, Any]) -> BaseAuthToken:
"""Create token object from API response"""
pass
@@ -369,14 +348,10 @@ class BaseAuthenticator(ABC):
# Try country-aware call first
try:
token_data = self.settings_manager.load_token_data(
self.provider_name, self.country
)
token_data = self.settings_manager.load_token_data(self.provider_name, self.country)
except TypeError:
# Fallback for managers that don't support country parameter
logger.debug(
f"Settings manager doesn't support country parameter, using without"
)
logger.debug(f"Settings manager doesn't support country parameter, using without")
token_data = self.settings_manager.load_token_data(self.provider_name)
if token_data:
@@ -385,9 +360,7 @@ class BaseAuthenticator(ABC):
f"Successfully loaded existing session for {self.provider_name}{country_str}"
)
else:
logger.info(
f"No existing session found for {self.provider_name}{country_str}"
)
logger.info(f"No existing session found for {self.provider_name}{country_str}")
self._current_token = None
except Exception as e:
@@ -417,9 +390,7 @@ class BaseAuthenticator(ABC):
if success:
logger.debug(f"Saved session for {self.provider_name}{country_str}")
else:
logger.warning(
f"Failed to save session for {self.provider_name}{country_str}"
)
logger.warning(f"Failed to save session for {self.provider_name}{country_str}")
except Exception as e:
logger.error(f"Error saving session for {self.provider_name}: {e}")
@@ -434,9 +405,7 @@ class BaseAuthenticator(ABC):
# If we still don't have valid credentials, try fallback
if not self.credentials or not self.credentials.validate():
logger.info(
f"No valid user credentials found for {self.provider_name}, using fallback"
)
logger.info(f"No valid user credentials found for {self.provider_name}, using fallback")
self.credentials = self.get_fallback_credentials()
return self.credentials is not None and self.credentials.validate()
@@ -460,14 +429,8 @@ class BaseAuthenticator(ABC):
logger.debug(f"[{self.provider_name}{country_str}] No current token")
# 1. Return existing token if valid
if (
not force_refresh
and self._current_token
and not self._current_token.is_expired
):
logger.info(
f"[{self.provider_name}{country_str}] Using existing valid token"
)
if not force_refresh and self._current_token and not self._current_token.is_expired:
logger.info(f"[{self.provider_name}{country_str}] Using existing valid token")
return self._current_token
# 2. ENHANCED REFRESH LOGIC: Attempt refresh if we have a token with refresh capability
@@ -492,9 +455,7 @@ class BaseAuthenticator(ABC):
if should_attempt_refresh:
logger.info(f"[{self.provider_name}{country_str}] Attempting token refresh")
try:
refreshed_token = (
self._refresh_token()
) # Provider-specific implementation
refreshed_token = self._refresh_token() # Provider-specific implementation
logger.debug(
f"[{self.provider_name}{country_str}] Refresh result: {refreshed_token is not None}"
)
@@ -502,27 +463,21 @@ class BaseAuthenticator(ABC):
if refreshed_token:
self._current_token = refreshed_token
self._save_session()
logger.info(
f"[{self.provider_name}{country_str}] Token refresh successful"
)
logger.info(f"[{self.provider_name}{country_str}] Token refresh successful")
return self._current_token
else:
logger.debug(
f"[{self.provider_name}{country_str}] Refresh failed, falling back to full auth"
)
except Exception as e:
logger.warning(
f"[{self.provider_name}{country_str}] Token refresh failed: {e}"
)
logger.warning(f"[{self.provider_name}{country_str}] Token refresh failed: {e}")
# 3. Ensure we have credentials before attempting full authentication
if not self._ensure_credentials():
raise Exception(f"No valid credentials available for {self.provider_name}")
# 4. Perform full authentication
logger.info(
f"[{self.provider_name}{country_str}] Performing new authentication"
)
logger.info(f"[{self.provider_name}{country_str}] Performing new authentication")
token = self._perform_authentication()
self._current_token = token
self._save_session()
@@ -540,9 +495,7 @@ class BaseAuthenticator(ABC):
)
except TypeError:
# Fallback for managers that don't support country parameter
logger.debug(
f"Settings manager doesn't support country parameter, using without"
)
logger.debug(f"Settings manager doesn't support country parameter, using without")
success = self.settings_manager.save_provider_credentials(
self.provider_name, credentials
)
@@ -574,14 +527,10 @@ class BaseAuthenticator(ABC):
return success
else:
logger.debug(
f"No Kodi sync capability available for {self.provider_name}"
)
logger.debug(f"No Kodi sync capability available for {self.provider_name}")
return True
except Exception as e:
logger.error(
f"Error syncing credentials from Kodi for {self.provider_name}: {e}"
)
logger.error(f"Error syncing credentials from Kodi for {self.provider_name}: {e}")
return False
def get_credential_info(self) -> Dict[str, Any]:
@@ -606,14 +555,10 @@ class BaseAuthenticator(ABC):
self.provider_name, self.country
)
except TypeError:
manager_info = self.settings_manager.get_credential_info(
self.provider_name
)
manager_info = self.settings_manager.get_credential_info(self.provider_name)
base_info.update(manager_info)
except Exception as e:
logger.debug(
f"Could not get extended credential info for {self.provider_name}: {e}"
)
logger.debug(f"Could not get extended credential info for {self.provider_name}: {e}")
return base_info
@@ -632,16 +577,12 @@ class BaseAuthenticator(ABC):
if hasattr(self.settings_manager, "credential_manager"):
# Try country-aware call first
try:
success = (
self.settings_manager.credential_manager.delete_credentials(
self.provider_name, self.country
)
success = self.settings_manager.credential_manager.delete_credentials(
self.provider_name, self.country
)
except TypeError:
success = (
self.settings_manager.credential_manager.delete_credentials(
self.provider_name
)
success = self.settings_manager.credential_manager.delete_credentials(
self.provider_name
)
else:
logger.debug(
@@ -649,19 +590,13 @@ class BaseAuthenticator(ABC):
)
if success:
logger.info(
f"{self.provider_name}: Stored credentials cleared successfully"
)
logger.info(f"{self.provider_name}: Stored credentials cleared successfully")
self.invalidate_token()
else:
logger.warning(
f"{self.provider_name}: Failed to clear stored credentials"
)
logger.warning(f"{self.provider_name}: Failed to clear stored credentials")
return success
except Exception as e:
logger.error(
f"Error clearing stored credentials for {self.provider_name}: {e}"
)
logger.error(f"Error clearing stored credentials for {self.provider_name}: {e}")
return False
def has_stored_credentials(self) -> bool:
@@ -673,14 +608,10 @@ class BaseAuthenticator(ABC):
self.provider_name, self.country
)
except TypeError:
credentials = self.settings_manager.get_provider_credentials(
self.provider_name
)
credentials = self.settings_manager.get_provider_credentials(self.provider_name)
return credentials is not None and credentials.validate()
except Exception as e:
logger.debug(
f"Error checking stored credentials for {self.provider_name}: {e}"
)
logger.debug(f"Error checking stored credentials for {self.provider_name}: {e}")
return False
def test_current_credentials(self) -> bool:
@@ -731,9 +662,7 @@ class BaseAuthenticator(ABC):
try:
# Try country-aware call first
try:
return self.settings_manager.get_device_id(
self.provider_name, self.country
)
return self.settings_manager.get_device_id(self.provider_name, self.country)
except TypeError:
return self.settings_manager.get_device_id(self.provider_name)
except Exception as e:
@@ -818,13 +747,10 @@ class BaseAuthenticator(ABC):
self.provider_name, self.country
)
except TypeError:
stored_creds = self.settings_manager.get_provider_credentials(
self.provider_name
)
stored_creds = self.settings_manager.get_provider_credentials(self.provider_name)
has_user_creds = (
isinstance(stored_creds, UserPasswordCredentials)
and stored_creds.validate()
isinstance(stored_creds, UserPasswordCredentials) and stored_creds.validate()
)
# Check current credentials
@@ -17,9 +17,7 @@ from .base_auth import BaseAuthenticator, BaseAuthToken, TokenAuthLevel
class OAuth2Error(Exception):
"""OAuth2-specific error with structured error information"""
def __init__(
self, error: str, error_description: str = None, error_uri: str = None
):
def __init__(self, error: str, error_description: str = None, error_uri: str = None):
self.error = error
self.error_description = error_description
self.error_uri = error_uri
@@ -117,9 +115,7 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
if self._http_manager is not None:
return self._http_manager
logger.warning(
f"No HTTP manager available for {self.provider_name}, creating one"
)
logger.warning(f"No HTTP manager available for {self.provider_name}, creating one")
try:
from ...base.network import HTTPManagerFactory
@@ -131,9 +127,7 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
timeout=getattr(self.config, "timeout", 30),
)
except Exception as e:
logger.warning(
f"Error creating HTTP manager via factory: {e}, using minimal fallback"
)
logger.warning(f"Error creating HTTP manager via factory: {e}, using minimal fallback")
self._http_manager = self._create_minimal_http_manager()
return self._http_manager
@@ -152,18 +146,14 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
return self._config
# Log who's calling this before config is set
logger.warning(
f"Config accessed before initialization for {self.provider_name}"
)
logger.warning(f"Config accessed before initialization for {self.provider_name}")
logger.debug(f"Call stack:\n{''.join(traceback.format_stack()[-5:])}")
# Only create minimal config if absolutely necessary
class MinimalConfig:
def __init__(self):
self.timeout = 30
self.user_agent = (
"Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"
)
self.user_agent = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"
self.base_website = "https://example.com"
self.auth_endpoint = "https://auth.example.com"
@@ -215,9 +205,7 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
if hasattr(self.config, "auth_endpoint"):
return self.config.auth_endpoint
raise NotImplementedError(
"Subclass must implement auth_endpoint or set _auth_endpoint"
)
raise NotImplementedError("Subclass must implement auth_endpoint or set _auth_endpoint")
@auth_endpoint.setter
def auth_endpoint(self, value):
@@ -246,9 +234,7 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
if hasattr(self, "auth_endpoint"):
auth_endpoint = self.auth_endpoint
else:
logger.warning(
f"auth_endpoint not defined for {self.provider_name}, using default"
)
logger.warning(f"auth_endpoint not defined for {self.provider_name}, using default")
return "https://auth.example.com/oauth2/auth"
if auth_endpoint.endswith("/token"):
@@ -320,9 +306,7 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
Uses provider-specific headers and payload formatting
"""
try:
logger.debug(
f"Starting OAuth2 client credentials flow for {self.provider_name}"
)
logger.debug(f"Starting OAuth2 client credentials flow for {self.provider_name}")
headers = self._get_auth_headers()
data = self._build_auth_payload()
@@ -335,23 +319,17 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
response.raise_for_status()
token_data = response.json()
logger.debug(
f"OAuth2 client credentials flow successful for {self.provider_name}"
)
logger.debug(f"OAuth2 client credentials flow successful for {self.provider_name}")
return token_data
except OAuth2Error:
raise
except Exception as e:
logger.error(
f"OAuth2 client credentials flow failed for {self.provider_name}: {e}"
)
logger.error(f"OAuth2 client credentials flow failed for {self.provider_name}: {e}")
raise Exception(f"OAuth2 client credentials flow failed: {e}")
# Authorization URL Building
def _build_authorization_url(
self, extra_params: Dict[str, Any] = None
) -> tuple[str, str, str]:
def _build_authorization_url(self, extra_params: Dict[str, Any] = None) -> tuple[str, str, str]:
"""Build authorization URL with PKCE for authorization code flow"""
code_verifier = self.generate_pkce_verifier()
code_challenge = self.generate_pkce_challenge(code_verifier)
@@ -383,9 +361,7 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
Enhanced to support provider-specific customizations
"""
try:
logger.debug(
f"Exchanging authorization code for token for {self.provider_name}"
)
logger.debug(f"Exchanging authorization code for token for {self.provider_name}")
# Allow subclasses to override the default payload
data = self._build_token_exchange_payload(
@@ -419,17 +395,13 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
response.raise_for_status()
token_data = response.json()
logger.debug(
f"Authorization code exchange successful for {self.provider_name}"
)
logger.debug(f"Authorization code exchange successful for {self.provider_name}")
return token_data
except OAuth2Error:
raise
except Exception as e:
logger.error(
f"Authorization code exchange failed for {self.provider_name}: {e}"
)
logger.error(f"Authorization code exchange failed for {self.provider_name}: {e}")
raise Exception(f"Authorization code exchange failed: {e}")
# New flexible methods that subclasses can override
@@ -456,9 +428,7 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
headers = self._get_auth_headers()
# Ensure Content-Type is appropriate
if kwargs.get("use_json", False) or self._should_use_json_for_token_exchange(
**kwargs
):
if kwargs.get("use_json", False) or self._should_use_json_for_token_exchange(**kwargs):
headers["Content-Type"] = "application/json"
else:
headers["Content-Type"] = "application/x-www-form-urlencoded"
@@ -510,9 +480,7 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
# Step 2: Extract login form action URL
form_matches = re.findall(form_selector_pattern, auth_response.text)
if not form_matches:
raise Exception(
f"Could not find login form using pattern: {form_selector_pattern}"
)
raise Exception(f"Could not find login form using pattern: {form_selector_pattern}")
login_url = html.unescape(form_matches[0])
@@ -538,19 +506,15 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
if not redirect_url:
raise Exception("No redirect URL found after login")
else:
redirect_response = session.get(
login_response.url, timeout=self.config.timeout
)
redirect_response = session.get(login_response.url, timeout=self.config.timeout)
redirect_url = redirect_response.url
# Step 6: Validate and extract authorization code
is_valid, error_msg, authorization_code = (
self.validate_authentication_response(redirect_url, state)
is_valid, error_msg, authorization_code = self.validate_authentication_response(
redirect_url, state
)
if not is_valid:
raise Exception(
f"Authentication response validation failed: {error_msg}"
)
raise Exception(f"Authentication response validation failed: {error_msg}")
# Step 7: Exchange code for token
token_data = self._exchange_authorization_code_for_token(
@@ -645,30 +609,22 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
try:
headers = self.config.get_base_headers()
response = self.http_manager.get(
main_page_url, operation="api", headers=headers
)
response = self.http_manager.get(main_page_url, operation="api", headers=headers)
response.raise_for_status()
js_matches = re.findall(js_file_pattern, response.text)
if not js_matches:
logger.warning(
f"Could not find JS file using pattern: {js_file_pattern}"
)
logger.warning(f"Could not find JS file using pattern: {js_file_pattern}")
return None
js_url = main_page_url.rstrip("/") + "/" + js_matches[-1].lstrip("/")
js_response = self.http_manager.get(
js_url, operation="api", headers=headers
)
js_response = self.http_manager.get(js_url, operation="api", headers=headers)
js_response.raise_for_status()
client_id_match = re.search(client_id_pattern, js_response.text)
if not client_id_match:
logger.warning(
f"Could not find client ID using pattern: {client_id_pattern}"
)
logger.warning(f"Could not find client ID using pattern: {client_id_pattern}")
return None
return client_id_match.group(1)
@@ -700,9 +656,7 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
try:
headers = self.config.get_base_headers()
response = self.http_manager.get(
main_page_url, operation="api", headers=headers
)
response = self.http_manager.get(main_page_url, operation="api", headers=headers)
response.raise_for_status()
js_matches = re.findall(js_file_pattern, response.text)
@@ -711,9 +665,7 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
js_url = main_page_url.rstrip("/") + "/" + js_matches[-1].lstrip("/")
js_response = self.http_manager.get(
js_url, operation="api", headers=headers
)
js_response = self.http_manager.get(js_url, operation="api", headers=headers)
js_response.raise_for_status()
config_match = re.search(config_pattern, js_response.text)
@@ -746,9 +698,7 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
from ...base.auth.credentials import UserPasswordCredentials
# ALWAYS check stored credentials first for user credentials
stored_creds = self.settings_manager.get_provider_credentials(
self.provider_name
)
stored_creds = self.settings_manager.get_provider_credentials(self.provider_name)
if stored_creds and isinstance(stored_creds, UserPasswordCredentials):
if stored_creds.validate():
return stored_creds
@@ -762,9 +712,7 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
self.credentials = fallback
return fallback
def get_bearer_token(
self, force_refresh: bool = False, force_upgrade: bool = False
) -> str:
def get_bearer_token(self, force_refresh: bool = False, force_upgrade: bool = False) -> str:
"""
Get bearer token with automatic upgrade support
@@ -788,9 +736,7 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
logger.debug(f"Token classified as: {current_token.auth_level.value}")
# Check if upgrade is needed/requested
should_upgrade = force_upgrade or self._should_upgrade_to_user_token(
current_token
)
should_upgrade = force_upgrade or self._should_upgrade_to_user_token(current_token)
if should_upgrade and not force_refresh:
logger.info(
@@ -841,8 +787,7 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
# Main Authentication Flow
def _perform_authentication(self) -> BaseAuthToken:
"""Complete OAuth2 authentication based on credential type"""
from ...base.auth.credentials import (ClientCredentials,
UserPasswordCredentials)
from ...base.auth.credentials import ClientCredentials, UserPasswordCredentials
logger.debug(
f"Starting OAuth2 authentication for {self.provider_name} with credential type: {type(self.credentials)}"
@@ -852,9 +797,7 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
try:
if isinstance(self.credentials, UserPasswordCredentials):
logger.info(
f"Attempting OAuth2 user authentication for {self.provider_name}"
)
logger.info(f"Attempting OAuth2 user authentication for {self.provider_name}")
token_data = self._perform_oauth_authorization_code_flow(
self.credentials.username, self.credentials.password
)
@@ -864,18 +807,14 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
)
token_data = self._perform_oauth_client_credentials_flow()
else:
raise Exception(
f"Unsupported credential type for OAuth2: {type(self.credentials)}"
)
raise Exception(f"Unsupported credential type for OAuth2: {type(self.credentials)}")
token = self._create_token_from_response(token_data)
logger.info(f"OAuth2 authentication successful for {self.provider_name}")
return token
except Exception as e:
logger.error(
f"Primary OAuth2 authentication failed for {self.provider_name}: {e}"
)
logger.error(f"Primary OAuth2 authentication failed for {self.provider_name}: {e}")
if isinstance(original_credentials, UserPasswordCredentials):
logger.info(
@@ -900,9 +839,7 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
# Token Management
@abstractmethod
def _create_token_from_response(
self, response_data: Dict[str, Any]
) -> BaseAuthToken:
def _create_token_from_response(self, response_data: Dict[str, Any]) -> BaseAuthToken:
"""Create provider-specific token from OAuth2 response"""
pass
@@ -924,9 +861,7 @@ class BaseOAuth2Authenticator(BaseAuthenticator):
"pkce_support": True,
"proxy_support": hasattr(self, "http_manager"),
"credential_type": type(self.credentials).__name__,
"has_refresh_token": bool(
self._current_token and self._current_token.refresh_token
),
"has_refresh_token": bool(self._current_token and self._current_token.refresh_token),
}
if self._current_token:
@@ -4,8 +4,7 @@ from typing import Any, Dict, List, Optional
# Import centralized logger and VFS
from ..utils.logger import logger
from .credentials import (BaseCredentials, ClientCredentials,
UserPasswordCredentials)
from .credentials import BaseCredentials, ClientCredentials, UserPasswordCredentials
class CredentialManager:
@@ -25,9 +24,7 @@ class CredentialManager:
# Ensure base directory exists
self.vfs.mkdirs("")
logger.debug(
f"CredentialManager initialized with VFS base: {self.vfs.base_path}"
)
logger.debug(f"CredentialManager initialized with VFS base: {self.vfs.base_path}")
@staticmethod
def _get_credential_path(provider: str, country: Optional[str] = None) -> tuple:
@@ -65,9 +62,7 @@ class CredentialManager:
country_str = f" (country: {country})" if country else ""
try:
logger.debug(
f"CredentialManager: Loading credentials for '{provider}{country_str}'"
)
logger.debug(f"CredentialManager: Loading credentials for '{provider}{country_str}'")
# Check if credentials file exists
if not self.vfs.exists(self.credentials_file):
@@ -76,20 +71,14 @@ class CredentialManager:
)
return None
logger.debug(
f"CredentialManager: Reading credentials file: {self.credentials_file}"
)
logger.debug(f"CredentialManager: Reading credentials file: {self.credentials_file}")
data = self.vfs.read_json(self.credentials_file)
if not data:
logger.debug(
f"CredentialManager: Credentials file is empty or invalid JSON"
)
logger.debug(f"CredentialManager: Credentials file is empty or invalid JSON")
return None
logger.debug(
f"CredentialManager: Available providers in file: {list(data.keys())}"
)
logger.debug(f"CredentialManager: Available providers in file: {list(data.keys())}")
# Try with country first (new format: {"provider": {"country": {...}}})
if country:
@@ -107,9 +96,7 @@ class CredentialManager:
logger.debug(
f"CredentialManager: Found credentials with country nesting for '{provider}{country_str}'"
)
return self._create_credential_from_data(
provider, country, provider_data
)
return self._create_credential_from_data(provider, country, provider_data)
# Fallback: try without country (legacy format: {"provider": {...}})
logger.debug(
@@ -130,9 +117,7 @@ class CredentialManager:
logger.debug(
f"CredentialManager: Found credentials in legacy format for '{provider}{country_str}'"
)
return self._create_credential_from_data(
provider, country, provider_data
)
return self._create_credential_from_data(provider, country, provider_data)
return None
@@ -149,9 +134,7 @@ class CredentialManager:
provider_data = provider_data[key]
if not provider_data:
logger.debug(
f"CredentialManager: No data found for provider '{provider}'"
)
logger.debug(f"CredentialManager: No data found for provider '{provider}'")
return None
return self._create_credential_from_data(provider, None, provider_data)
@@ -224,9 +207,7 @@ class CredentialManager:
creds = ClientCredentials(
client_id=client_id,
client_secret=self._decode_password(
provider_data.get("client_secret", "")
),
client_secret=self._decode_password(provider_data.get("client_secret", "")),
grant_type=provider_data.get("grant_type", "client_credentials"),
)
@@ -308,9 +289,7 @@ class CredentialManager:
success = self.vfs.write_json(self.credentials_file, data)
if success:
logger.info(
f"Successfully saved credentials for {provider}{country_str}"
)
logger.info(f"Successfully saved credentials for {provider}{country_str}")
# Verify by reading back
verify_data = self.vfs.read_json(self.credentials_file)
@@ -318,10 +297,7 @@ class CredentialManager:
verify_current = verify_data
found = True
for key in keys_path:
if (
not isinstance(verify_current, dict)
or key not in verify_current
):
if not isinstance(verify_current, dict) or key not in verify_current:
found = False
break
verify_current = verify_current[key]
@@ -381,16 +357,12 @@ class CredentialManager:
and country in data[provider]
):
del data[provider][country]
logger.info(
f"Deleted stored credentials for {provider}{country_str}"
)
logger.info(f"Deleted stored credentials for {provider}{country_str}")
# If provider dict is now empty, remove it entirely
if not data[provider]:
del data[provider]
logger.debug(
f"Provider {provider} had no more countries, removed entirely"
)
logger.debug(f"Provider {provider} had no more countries, removed entirely")
return self.vfs.write_json(self.credentials_file, data)
else:
@@ -404,9 +376,7 @@ class CredentialManager:
logger.info(f"Deleted all stored credentials for {provider}")
return self.vfs.write_json(self.credentials_file, data)
else:
logger.debug(
f"No stored credentials found to delete for {provider}"
)
logger.debug(f"No stored credentials found to delete for {provider}")
return True
@@ -424,9 +394,7 @@ class CredentialManager:
"""
try:
if not self.vfs.exists(self.credentials_file):
logger.debug(
"No credentials file exists, returning empty provider list"
)
logger.debug("No credentials file exists, returning empty provider list")
return []
data = self.vfs.read_json(self.credentials_file)
@@ -549,9 +517,7 @@ class CredentialManager:
# Non-country provider
credentials = self.load_credentials(provider)
if credentials:
export_data["credentials"][provider] = self._credential_to_dict(
credentials
)
export_data["credentials"][provider] = self._credential_to_dict(credentials)
else:
# Export all providers
for prov in self.list_providers():
@@ -569,9 +535,7 @@ class CredentialManager:
# Non-country provider
credentials = self.load_credentials(prov)
if credentials:
export_data["credentials"][prov] = self._credential_to_dict(
credentials
)
export_data["credentials"][prov] = self._credential_to_dict(credentials)
return export_data
@@ -624,9 +588,7 @@ class CredentialManager:
credentials = self._dict_to_credential(cred_data)
if credentials:
success = self.save_credentials(
provider, credentials, country
)
success = self.save_credentials(provider, credentials, country)
results[f"{provider}_{country}"] = success
if success:
@@ -645,9 +607,7 @@ class CredentialManager:
results[provider] = success
if success:
logger.info(
f"Successfully imported credentials for {provider}"
)
logger.info(f"Successfully imported credentials for {provider}")
else:
logger.error(f"Failed to import credentials for {provider}")
@@ -686,9 +646,7 @@ class CredentialManager:
credentials_file_path = self.vfs.join_path(self.credentials_file)
logger.info(f"Credentials file path: {credentials_file_path}")
logger.info(
f"Credentials file exists: {self.vfs.exists(self.credentials_file)}"
)
logger.info(f"Credentials file exists: {self.vfs.exists(self.credentials_file)}")
if self.vfs.exists(self.credentials_file):
file_size = self.vfs.get_size(self.credentials_file)
@@ -712,10 +670,7 @@ class CredentialManager:
if has_countries:
logger.info(f" {provider} (country-aware):")
for country, cred_data in provider_data.items():
if (
isinstance(cred_data, dict)
and "type" in cred_data
):
if isinstance(cred_data, dict) and "type" in cred_data:
cred_type = cred_data.get("type")
username = cred_data.get("username", "N/A")
logger.info(
@@ -729,9 +684,7 @@ class CredentialManager:
f" {provider} (no country): type={cred_type}, username={username}"
)
else:
logger.error(
"Credentials file contains invalid JSON or is empty"
)
logger.error("Credentials file contains invalid JSON or is empty")
else:
logger.info("Credentials file is empty")
else:
@@ -85,22 +85,16 @@ class SessionManager:
for key in keys_path:
if not isinstance(session_data, dict) or key not in session_data:
logger.info(
f"No session data found at path: {' -> '.join(keys_path)}"
)
logger.info(f"No session data found at path: {' -> '.join(keys_path)}")
return None
session_data = session_data[key]
if session_data:
# Log what we found (without sensitive data)
safe_data = self._get_safe_representation(session_data)
logger.info(
f"Loaded session data for {provider}{country_str}: {safe_data}"
)
logger.info(f"Loaded session data for {provider}{country_str}: {safe_data}")
else:
logger.info(
f"No session data found for {provider}{country_str} in session file"
)
logger.info(f"No session data found for {provider}{country_str} in session file")
return session_data
@@ -173,9 +167,7 @@ class SessionManager:
success = self.vfs.write_json(self.session_file, data)
if success:
logger.info(
f"Successfully saved session data for {provider}{country_str}"
)
logger.info(f"Successfully saved session data for {provider}{country_str}")
# Verify by reading it back
verify_data = self.vfs.read_json(self.session_file)
@@ -183,10 +175,7 @@ class SessionManager:
verify_current = verify_data
found = True
for key in keys_path:
if (
not isinstance(verify_current, dict)
or key not in verify_current
):
if not isinstance(verify_current, dict) or key not in verify_current:
found = False
break
verify_current = verify_current[key]
@@ -201,9 +190,7 @@ class SessionManager:
)
return False
else:
logger.error(
f"Verification failed: Could not read back session file"
)
logger.error(f"Verification failed: Could not read back session file")
return False
else:
logger.error(f"Failed to write session file")
@@ -240,9 +227,7 @@ class SessionManager:
country_str = f" (country: {country})" if country else ""
try:
logger.debug(
f"Saving scoped token for {provider}{country_str}, scope: {scope}"
)
logger.debug(f"Saving scoped token for {provider}{country_str}, scope: {scope}")
# Load existing session data
session_data = self.load_session(provider, country) or {}
@@ -252,25 +237,17 @@ class SessionManager:
# Log what we're saving
safe_token_data = self._get_safe_representation(token_data)
logger.info(
f"Scoped token data for {provider}{country_str}/{scope}: {safe_token_data}"
)
logger.info(f"Scoped token data for {provider}{country_str}/{scope}: {safe_token_data}")
success = self.save_session(provider, session_data, country)
if success:
logger.info(
f"Successfully saved scoped token for {provider}{country_str}/{scope}"
)
logger.info(f"Successfully saved scoped token for {provider}{country_str}/{scope}")
else:
logger.error(
f"Failed to save scoped token for {provider}{country_str}/{scope}"
)
logger.error(f"Failed to save scoped token for {provider}{country_str}/{scope}")
return success
except Exception as e:
logger.error(
f"Error saving scoped token for {provider}{country_str}/{scope}: {e}"
)
logger.error(f"Error saving scoped token for {provider}{country_str}/{scope}: {e}")
return False
def load_scoped_token(
@@ -281,9 +258,7 @@ class SessionManager:
"""
country_str = f" (country: {country})" if country else ""
logger.debug(
f"Loading scoped token for {provider}{country_str}, scope: {scope}"
)
logger.debug(f"Loading scoped token for {provider}{country_str}, scope: {scope}")
session_data = self.load_session(provider, country)
if not session_data:
@@ -292,9 +267,7 @@ class SessionManager:
# Check if scope exists
if scope not in session_data:
logger.info(
f"No token found for scope '{scope}' in {provider}{country_str}"
)
logger.info(f"No token found for scope '{scope}' in {provider}{country_str}")
logger.debug(f"Available scopes: {list(session_data.keys())}")
return None
@@ -303,9 +276,7 @@ class SessionManager:
# 🚨 FIX: Special handling for persona scope
if scope == "persona":
if not isinstance(token_data, dict) or "persona_token" not in token_data:
logger.warning(
f"Persona scope exists but doesn't contain valid persona token data"
)
logger.warning(f"Persona scope exists but doesn't contain valid persona token data")
return None
# Persona tokens don't expire in the same way - they use expires_at field
@@ -314,9 +285,7 @@ class SessionManager:
# Validate it's actually token data (for other scopes)
if not isinstance(token_data, dict) or "access_token" not in token_data:
logger.warning(
f"Scope '{scope}' exists but doesn't contain valid token data"
)
logger.warning(f"Scope '{scope}' exists but doesn't contain valid token data")
return None
# Check access token expiration
@@ -332,9 +301,7 @@ class SessionManager:
)
return token_data # Return so caller can attempt refresh
elif access_token_expired and refresh_token_expired:
logger.warning(
f"Both access and refresh tokens expired for scope '{scope}'"
)
logger.warning(f"Both access and refresh tokens expired for scope '{scope}'")
return None # Both expired, no point returning
else:
logger.info(f"Loaded valid token for {provider}{country_str}/{scope}")
@@ -350,9 +317,7 @@ class SessionManager:
return token_data
@staticmethod
def _is_token_expired(
token_data: Dict[str, Any], buffer_seconds: int = 300
) -> bool:
def _is_token_expired(token_data: Dict[str, Any], buffer_seconds: int = 300) -> bool:
"""
Check if token is expired with buffer
@@ -393,10 +358,7 @@ class SessionManager:
# Format 2: yo_digital separate expiration for access_token
# yo_digital tokens have separate expiry for access and refresh tokens
if (
"access_token_expires_in" in token_data
and "access_token_issued_at" in token_data
):
if "access_token_expires_in" in token_data and "access_token_issued_at" in token_data:
expires_in = token_data.get("access_token_expires_in", 0)
issued_at = token_data.get("access_token_issued_at", 0)
expires_at = issued_at + expires_in
@@ -430,9 +392,7 @@ class SessionManager:
return False
@staticmethod
def _is_refresh_token_expired(
token_data: Dict[str, Any], buffer_seconds: int = 300
) -> bool:
def _is_refresh_token_expired(token_data: Dict[str, Any], buffer_seconds: int = 300) -> bool:
"""
Check if refresh token is expired (for yo_digital tokens)
@@ -447,10 +407,7 @@ class SessionManager:
True if refresh token expired, False otherwise or if no refresh token
"""
# yo_digital format with separate refresh token expiry
if (
"refresh_token_expires_in" in token_data
and "refresh_token_issued_at" in token_data
):
if "refresh_token_expires_in" in token_data and "refresh_token_issued_at" in token_data:
current_time = time.time()
expires_in = token_data.get("refresh_token_expires_in", 0)
issued_at = token_data.get("refresh_token_issued_at", 0)
@@ -497,21 +454,15 @@ class SessionManager:
# Log token info (without sensitive data)
safe_token_data = self._get_safe_representation(token_data)
logger.info(
f"Token data to save for {provider}{country_str}: {safe_token_data}"
)
logger.info(f"Token data to save for {provider}{country_str}: {safe_token_data}")
session_data.update(token_data)
success = self.save_session(provider, session_data, country)
if success:
logger.info(
f"Successfully saved authentication token for {provider}{country_str}"
)
logger.info(f"Successfully saved authentication token for {provider}{country_str}")
else:
logger.error(
f"Failed to save authentication token for {provider}{country_str}"
)
logger.error(f"Failed to save authentication token for {provider}{country_str}")
return success
except Exception as e:
@@ -528,16 +479,12 @@ class SessionManager:
session_data = self.load_session(provider, country)
if not session_data:
logger.info(
f"No session data available for token loading for {provider}{country_str}"
)
logger.info(f"No session data available for token loading for {provider}{country_str}")
return None
# Check if we have token data
if "access_token" not in session_data:
logger.info(
f"No access token found in session data for {provider}{country_str}"
)
logger.info(f"No access token found in session data for {provider}{country_str}")
logger.debug(f"Available session keys: {list(session_data.keys())}")
return None
@@ -569,13 +516,9 @@ class SessionManager:
device_id = str(uuid.uuid4())
session_data["device_id"] = device_id
self.save_session(provider, session_data, country)
logger.info(
f"Generated new device ID for {provider}{country_str}: {device_id}"
)
logger.info(f"Generated new device ID for {provider}{country_str}: {device_id}")
else:
logger.debug(
f"Using existing device ID for {provider}{country_str}: {device_id}"
)
logger.debug(f"Using existing device ID for {provider}{country_str}: {device_id}")
return device_id
@@ -602,15 +545,11 @@ class SessionManager:
if not data[provider]:
del data[provider]
logger.debug(
f"Provider {provider} had no more countries, removed entirely"
)
logger.debug(f"Provider {provider} had no more countries, removed entirely")
return self.vfs.write_json(self.session_file, data)
else:
logger.debug(
f"No session data found to clear for {provider}{country_str}"
)
logger.debug(f"No session data found to clear for {provider}{country_str}")
else:
if provider in data:
del data[provider]
@@ -625,9 +564,7 @@ class SessionManager:
logger.error(f"Error clearing session for {provider}{country_str}: {e}")
return False
def clear_scoped_token(
self, provider: str, scope: str, country: Optional[str] = None
) -> bool:
def clear_scoped_token(self, provider: str, scope: str, country: Optional[str] = None) -> bool:
"""
Clear token for a specific scope
@@ -652,15 +589,11 @@ class SessionManager:
logger.info(f"Cleared scoped token for {provider}{country_str}/{scope}")
return self.save_session(provider, session_data, country)
else:
logger.debug(
f"No token found for scope '{scope}' in {provider}{country_str}"
)
logger.debug(f"No token found for scope '{scope}' in {provider}{country_str}")
return True
except Exception as e:
logger.error(
f"Error clearing scoped token for {provider}{country_str}/{scope}: {e}"
)
logger.error(f"Error clearing scoped token for {provider}{country_str}/{scope}: {e}")
return False
def clear_token(self, provider: str, country: Optional[str] = None) -> bool:
@@ -670,9 +603,7 @@ class SessionManager:
try:
session_data = self.load_session(provider, country)
if not session_data:
logger.debug(
f"No session data found, nothing to clear for {provider}{country_str}"
)
logger.debug(f"No session data found, nothing to clear for {provider}{country_str}")
return True
# Remove token-related fields
@@ -691,9 +622,7 @@ class SessionManager:
fields_removed.append(field)
if fields_removed:
logger.debug(
f"Cleared token fields {fields_removed} for {provider}{country_str}"
)
logger.debug(f"Cleared token fields {fields_removed} for {provider}{country_str}")
return self.save_session(provider, session_data, country)
@@ -783,17 +712,11 @@ class SessionManager:
logger.info(f" {provider} (country-aware):")
for country, session_data in provider_data.items():
if isinstance(session_data, dict):
safe_repr = self._get_safe_representation(
session_data
)
safe_repr = self._get_safe_representation(session_data)
logger.info(f" {country}: {safe_repr}")
else:
safe_repr = self._get_safe_representation(
provider_data
)
logger.info(
f" {provider} (no country): {safe_repr}"
)
safe_repr = self._get_safe_representation(provider_data)
logger.info(f" {provider} (no country): {safe_repr}")
else:
logger.error("Session file contains invalid JSON or is empty")
else:
@@ -84,9 +84,7 @@ class CatchupOperations:
provider_name, channel_id, country=country
)
def get_catchup_window(
self, provider_name: str, channel_id: Optional[str] = None
) -> int:
def get_catchup_window(self, provider_name: str, channel_id: Optional[str] = None) -> int:
"""Get catchup window in hours."""
provider = self.registry.get_provider(provider_name)
if not provider:
@@ -116,8 +114,7 @@ class CatchupOperations:
capabilities[name] = {
"supports_catchup": provider.supports_catchup,
"catchup_window": provider.catchup_window,
"catchup_enabled": provider.supports_catchup
and provider.catchup_window > 0,
"catchup_enabled": provider.supports_catchup and provider.catchup_window > 0,
}
except Exception as e:
logger.warning(f"Error getting catchup for '{name}': {e}")
@@ -39,9 +39,7 @@ class ChannelOperations:
return channels
def get_channel_manifest(
self, provider_name: str, channel_id: str, **kwargs
) -> Optional[str]:
def get_channel_manifest(self, provider_name: str, channel_id: str, **kwargs) -> Optional[str]:
"""Get manifest URL for a specific channel."""
provider = self.registry.get_provider(provider_name)
if not provider:
@@ -49,9 +47,7 @@ class ChannelOperations:
manifest_url = provider.get_manifest(channel_id, **kwargs)
if manifest_url:
logger.debug(
f"Retrieved manifest for '{channel_id}' from '{provider_name}'"
)
logger.debug(f"Retrieved manifest for '{channel_id}' from '{provider_name}'")
return manifest_url
def get_all_channels(
@@ -20,18 +20,14 @@ class DRMPluginManager:
logger.debug("DRMPluginManager: Initialized with empty plugin registry")
if auto_discover:
logger.debug(
"DRMPluginManager: Auto-discovery enabled, discovering plugins..."
)
logger.debug("DRMPluginManager: Auto-discovery enabled, discovering plugins...")
discovered = self.discover_plugins()
if discovered:
logger.debug(
f"DRMPluginManager: Auto-discovery completed, {len(discovered)} plugins ready"
)
else:
logger.debug(
"DRMPluginManager: Auto-discovery completed, no plugins found"
)
logger.debug("DRMPluginManager: Auto-discovery completed, no plugins found")
def register_plugin(self, plugin: DRMPlugin) -> None:
"""
@@ -88,9 +84,7 @@ class DRMPluginManager:
# Check if plugins directory exists
if not os.path.exists(plugins_dir):
logger.debug(
f"DRMPluginManager: Plugins directory does not exist: {plugins_dir}"
)
logger.debug(f"DRMPluginManager: Plugins directory does not exist: {plugins_dir}")
return []
logger.debug(f"DRMPluginManager: Scanning plugins directory: {plugins_dir}")
@@ -125,9 +119,7 @@ class DRMPluginManager:
module_name = os.path.splitext(filename)[0]
# Import the module dynamically
spec = importlib.util.spec_from_file_location(
module_name, file_path
)
spec = importlib.util.spec_from_file_location(module_name, file_path)
if spec is None or spec.loader is None:
logger.debug(
f"DRMPluginManager: Could not create module spec for {filename}"
@@ -149,9 +141,7 @@ class DRMPluginManager:
plugin_classes.append((name, obj))
if not plugin_classes:
logger.debug(
f"DRMPluginManager: No DRMPlugin classes found in {filename}"
)
logger.debug(f"DRMPluginManager: No DRMPlugin classes found in {filename}")
continue
logger.debug(
@@ -182,17 +172,15 @@ class DRMPluginManager:
)
except Exception as e:
error_msg = f"Failed to instantiate {class_name} from {filename}: {str(e)}"
failed_plugins.append(
(f"{filename}::{class_name}", error_msg)
error_msg = (
f"Failed to instantiate {class_name} from {filename}: {str(e)}"
)
failed_plugins.append((f"{filename}::{class_name}", error_msg))
logger.warning(f"DRMPluginManager: {error_msg}")
except Exception as e:
traceback_str = traceback.format_exc()
error_msg = (
f"Failed to process file {filename}: {str(e)}\n{traceback_str}"
)
error_msg = f"Failed to process file {filename}: {str(e)}\n{traceback_str}"
failed_plugins.append((filename, error_msg))
logger.warning(f"DRMPluginManager: {error_msg}")
@@ -214,9 +202,7 @@ class DRMPluginManager:
)
if failed_plugins:
logger.debug(
f"DRMPluginManager: {len(failed_plugins)} plugins/files failed to load:"
)
logger.debug(f"DRMPluginManager: {len(failed_plugins)} plugins/files failed to load:")
for plugin_name, error in failed_plugins:
logger.debug(f" - {plugin_name}: {error}")
@@ -262,9 +248,7 @@ class DRMPluginManager:
for drm_system, plugin in self.plugins.items():
if drm_system == DRMSystem.GENERIC:
generic_plugins.append(plugin)
logger.debug(
f"DRMPluginManager: Found generic plugin '{plugin.plugin_name}'"
)
logger.debug(f"DRMPluginManager: Found generic plugin '{plugin.plugin_name}'")
else:
specific_plugins.append((drm_system, plugin))
@@ -303,9 +287,7 @@ class DRMPluginManager:
# Process specific plugins
final_configs = []
for config in processed_configs:
logger.debug(
f"DRMPluginManager: Processing DRM config for system: {config.system}"
)
logger.debug(f"DRMPluginManager: Processing DRM config for system: {config.system}")
# Find specific plugin for this DRM system
plugin = self.plugins.get(config.system)
@@ -326,9 +308,7 @@ class DRMPluginManager:
f"DRMPluginManager: No PSSH data available for DRM system {config.system}"
)
processed_config = plugin.process_drm_config(
config, pssh_data, **kwargs
)
processed_config = plugin.process_drm_config(config, pssh_data, **kwargs)
if processed_config is not None:
# Check for ClearKey and return immediately if found
if processed_config.system == DRMSystem.CLEARKEY:
@@ -379,9 +359,7 @@ class DRMPluginManager:
f"DRMPluginManager: Retrieved plugin '{plugin.plugin_name}' for DRM system {drm_system}"
)
else:
logger.debug(
f"DRMPluginManager: No plugin found for DRM system {drm_system}"
)
logger.debug(f"DRMPluginManager: No plugin found for DRM system {drm_system}")
return plugin
def list_plugins(self) -> Dict[DRMSystem, str]:
@@ -392,8 +370,7 @@ class DRMPluginManager:
Dictionary mapping DRM systems to plugin names
"""
plugin_list = {
drm_system: plugin.plugin_name
for drm_system, plugin in self.plugins.items()
drm_system: plugin.plugin_name for drm_system, plugin in self.plugins.items()
}
logger.debug(f"DRMPluginManager: Currently registered plugins: {plugin_list}")
return plugin_list
+92 -15
View File
@@ -1,27 +1,62 @@
# ============================================================================
# streaming_providers/base/drm_operations.py
"""
DRM-related operations.
DRM-related operations with caching and optimized PSSH extraction.
"""
from typing import Dict, List
import time
from threading import Lock
from typing import Dict, List, Optional, Tuple
from .drm import DRMPluginManager
from .models import DRMSystem
from .utils.logger import logger
class PSSHCache:
"""Thread-safe cache for PSSH data"""
def __init__(self, ttl_seconds: int = 3600):
self.cache: Dict[str, Tuple[List, float]] = {}
self.ttl = ttl_seconds
self.lock = Lock()
def get(self, key: str) -> Optional[List]:
"""Get cached PSSH data if not expired"""
with self.lock:
if key in self.cache:
pssh_list, timestamp = self.cache[key]
if time.time() - timestamp < self.ttl:
logger.debug(f"Cache HIT for {key}")
return pssh_list
else:
logger.debug(f"Cache EXPIRED for {key}")
del self.cache[key]
return None
def set(self, key: str, pssh_list: List):
"""Cache PSSH data"""
with self.lock:
self.cache[key] = (pssh_list, time.time())
logger.debug(f"Cache SET for {key}")
def clear(self):
"""Clear all cache entries"""
with self.lock:
self.cache.clear()
logger.debug("Cache CLEARED")
class DRMOperations:
"""Handles all DRM-related operations."""
def __init__(self, registry):
def __init__(self, registry, cache_ttl: int = 3600):
self.registry = registry
self.drm_plugin_manager = DRMPluginManager()
logger.debug("DRMOperations: Initialized")
self.pssh_cache = PSSHCache(ttl_seconds=cache_ttl)
logger.debug("DRMOperations: Initialized with caching")
def get_channel_drm_configs(
self, provider_name: str, channel_id: str, **kwargs
) -> List:
def get_channel_drm_configs(self, provider_name: str, channel_id: str, **kwargs) -> List:
"""Get DRM configurations for a channel."""
provider = self.registry.get_provider(provider_name)
if not provider:
@@ -35,7 +70,15 @@ class DRMOperations:
if self._needs_pssh_extraction(drm_configs):
manifest_url = provider.get_manifest(channel_id, **kwargs)
if manifest_url:
pssh_data_list = self._extract_pssh_from_manifest(manifest_url)
# Try cache first
cache_key = f"{provider_name}:{channel_id}"
pssh_data_list = self.pssh_cache.get(cache_key)
if pssh_data_list is None:
# Cache miss - extract and cache
pssh_data_list = self._extract_pssh_from_manifest(manifest_url)
if pssh_data_list:
self.pssh_cache.set(cache_key, pssh_data_list)
# Process through plugins
processed = self.drm_plugin_manager.process_drm_configs(
@@ -50,12 +93,11 @@ class DRMOperations:
config_systems = {config.system for config in drm_configs}
plugin_systems = set(self.drm_plugin_manager.plugins.keys())
return bool(
config_systems & plugin_systems or DRMSystem.GENERIC in plugin_systems
)
return bool(config_systems & plugin_systems or DRMSystem.GENERIC in plugin_systems)
def _extract_pssh_from_manifest(self, manifest_url: str) -> List:
"""Extract PSSH data from manifest."""
@staticmethod
def _extract_pssh_from_manifest(manifest_url: str) -> List:
"""Extract PSSH data from manifest with single segment fallback."""
import requests
from .utils.manifest_parser import ManifestParser
@@ -63,9 +105,40 @@ class DRMOperations:
try:
response = requests.get(manifest_url, timeout=10)
response.raise_for_status()
return ManifestParser.extract_pssh_from_manifest(
response.text, manifest_url
manifest_content = response.text
# Extract from manifest content
pssh_list = ManifestParser._extract_from_manifest_content(manifest_content)
# Check if we need segment extraction
needs_segment_extraction = not pssh_list or any(
not p.pssh_box or not p.key_ids for p in pssh_list
)
if needs_segment_extraction:
logger.debug("PSSH incomplete in manifest, extracting from init segment")
# Extract ONE init segment URL
init_segment_url = ManifestParser.extract_single_init_segment_url(
manifest_content, manifest_url
)
if init_segment_url:
# Get expected system IDs
expected_system_ids = [p.system_id for p in pssh_list] if pssh_list else []
segment_pssh = ManifestParser._extract_from_single_segment(
init_segment_url, expected_system_ids
)
if segment_pssh:
# Merge or replace
return ManifestParser._merge_pssh_data(pssh_list, segment_pssh)
else:
logger.warning("Could not extract init segment URL from manifest")
return pssh_list
except Exception as e:
logger.warning(f"Failed to extract PSSH: {e}")
return []
@@ -77,3 +150,7 @@ class DRMOperations:
def clear_drm_plugins(self):
"""Clear all DRM plugins."""
self.drm_plugin_manager.clear_plugins()
def clear_pssh_cache(self):
"""Clear PSSH cache."""
self.pssh_cache.clear()
@@ -97,9 +97,7 @@ class EPGCache:
age = int(time.time()) - downloaded_at
if age > self.CACHE_TTL_SECONDS:
logger.info(
f"EPGCache: Cache expired (age: {age}s, TTL: {self.CACHE_TTL_SECONDS}s)"
)
logger.info(f"EPGCache: Cache expired (age: {age}s, TTL: {self.CACHE_TTL_SECONDS}s)")
return False
logger.debug(f"EPGCache: Cache valid (age: {age}s)")
@@ -148,11 +146,7 @@ class EPGCache:
# Determine if content is gzipped
content_type = response.headers.get("Content-Type", "").lower()
content_encoding = response.headers.get("Content-Encoding", "").lower()
is_gzipped = (
"gzip" in content_encoding
or url.endswith(".gz")
or "gzip" in content_type
)
is_gzipped = "gzip" in content_encoding or url.endswith(".gz") or "gzip" in content_type
filename = self.EPG_GZ_FILE if is_gzipped else self.EPG_FILE
+11 -37
View File
@@ -146,9 +146,7 @@ class EPGManager:
# Step 1: Map to EPG channel ID
epg_channel_id = self.mapping.get_epg_channel_id(provider_name, channel_id)
if not epg_channel_id:
logger.warning(
f"EPGManager: No EPG mapping found for {provider_name}/{channel_id}"
)
logger.warning(f"EPGManager: No EPG mapping found for {provider_name}/{channel_id}")
return []
logger.debug(f"EPGManager: Mapped to EPG channel ID: {epg_channel_id}")
@@ -218,9 +216,7 @@ class EPGManager:
try:
epg_channel_id = self.mapping.get_epg_channel_id(provider_name, channel_id)
if not epg_channel_id:
logger.warning(
f"EPGManager: No EPG mapping found for {provider_name}/{channel_id}"
)
logger.warning(f"EPGManager: No EPG mapping found for {provider_name}/{channel_id}")
return []
xml_path = self.cache.get_or_download(self.epg_url)
@@ -267,9 +263,7 @@ class EPGManager:
Returns:
Dictionary mapping channel IDs to their EPG entries (as dicts)
"""
logger.info(
f"EPGManager: Getting EPG for all channels of provider '{provider_name}'"
)
logger.info(f"EPGManager: Getting EPG for all channels of provider '{provider_name}'")
result = {}
@@ -277,9 +271,7 @@ class EPGManager:
provider_mapping = self.mapping.get_provider_mapping(provider_name)
if not provider_mapping:
logger.warning(
f"EPGManager: No channels mapped for provider '{provider_name}'"
)
logger.warning(f"EPGManager: No channels mapped for provider '{provider_name}'")
return result
# Get EPG for each channel
@@ -395,14 +387,8 @@ class EPGManager:
addon = xbmcaddon.Addon()
kodi_url = addon.getSetting("epg_xml_url")
if (
kodi_url
and kodi_url.strip()
and kodi_url != "https://example.com/epg.xml.gz"
):
logger.info(
f"EPGManager: Using EPG URL from Kodi settings: {kodi_url}"
)
if kodi_url and kodi_url.strip() and kodi_url != "https://example.com/epg.xml.gz":
logger.info(f"EPGManager: Using EPG URL from Kodi settings: {kodi_url}")
return kodi_url.strip()
except Exception as e:
logger.debug(f"EPGManager: Could not get EPG URL from Kodi settings: {e}")
@@ -411,24 +397,16 @@ class EPGManager:
try:
env_manager = get_environment_manager()
config_url = env_manager.get_config("epg_url")
if (
config_url
and config_url.strip()
and config_url != "https://example.com/epg.xml.gz"
):
if config_url and config_url.strip() and config_url != "https://example.com/epg.xml.gz":
logger.info(f"EPGManager: Using EPG URL from config.json: {config_url}")
return config_url.strip()
except Exception as e:
logger.debug(
f"EPGManager: Could not get EPG URL from environment manager: {e}"
)
logger.debug(f"EPGManager: Could not get EPG URL from environment manager: {e}")
# 4. Environment variable
env_url = os.environ.get("ULTIMATE_EPG_URL")
if env_url and env_url.strip() and env_url != "https://example.com/epg.xml.gz":
logger.info(
f"EPGManager: Using EPG URL from environment variable: {env_url}"
)
logger.info(f"EPGManager: Using EPG URL from environment variable: {env_url}")
return env_url.strip()
# 5. Try to get the last known URL from cache metadata
@@ -451,12 +429,8 @@ class EPGManager:
# 6. Default fallback (LAST RESORT - should rarely be used)
default_url = "https://example.com/epg.xml.gz"
logger.warning(
f"EPGManager: No valid EPG URL found, using default: {default_url}"
)
logger.warning(
"Please set ULTIMATE_EPG_URL environment variable or configure in settings!"
)
logger.warning(f"EPGManager: No valid EPG URL found, using default: {default_url}")
logger.warning("Please set ULTIMATE_EPG_URL environment variable or configure in settings!")
return default_url
@staticmethod
+15 -44
View File
@@ -50,9 +50,7 @@ class EPGMapping:
# Structure: {provider_name: {"mapping": {...}, "names": {...}}}
self._cache: Dict[str, Dict[str, Any]] = {}
logger.info(
f"EPGMapping: Initialized with user path: {self.user_vfs.base_path}"
)
logger.info(f"EPGMapping: Initialized with user path: {self.user_vfs.base_path}")
# Copy default mapping files from addon resources
self._copy_default_mapping_files()
@@ -74,19 +72,13 @@ class EPGMapping:
addon_info = bridge.get_addon_info()
addon_path = addon_info.get("path")
if addon_path:
default_dir = os.path.join(
addon_path, "resources", "config", "epg_mappings"
)
logger.debug(
f"EPGMapping: Default mapping dir (Kodi): {default_dir}"
)
default_dir = os.path.join(addon_path, "resources", "config", "epg_mappings")
logger.debug(f"EPGMapping: Default mapping dir (Kodi): {default_dir}")
return default_dir
# Fallback to standard filesystem
addon_path = os.getcwd()
default_dir = os.path.join(
addon_path, "resources", "config", "epg_mappings"
)
default_dir = os.path.join(addon_path, "resources", "config", "epg_mappings")
logger.debug(f"EPGMapping: Default mapping dir (standard): {default_dir}")
return default_dir
@@ -94,9 +86,7 @@ class EPGMapping:
logger.warning(f"EPGMapping: Could not determine default mapping dir: {e}")
# Last resort fallback
addon_path = os.getcwd()
default_dir = os.path.join(
addon_path, "resources", "config", "epg_mappings"
)
default_dir = os.path.join(addon_path, "resources", "config", "epg_mappings")
return default_dir
def _copy_default_mapping_files(self) -> bool:
@@ -114,9 +104,7 @@ class EPGMapping:
# Check if any user mapping files already exist
user_files = self.user_vfs.list_files(pattern=self.MAPPING_FILE_PATTERN)
if user_files:
logger.info(
f"EPGMapping: Found {len(user_files)} existing user mapping files"
)
logger.info(f"EPGMapping: Found {len(user_files)} existing user mapping files")
return True
else:
logger.warning("EPGMapping: No default files and no user files found")
@@ -124,14 +112,10 @@ class EPGMapping:
try:
# Get list of default mapping files
default_files = glob.glob(
os.path.join(default_dir, self.MAPPING_FILE_PATTERN)
)
default_files = glob.glob(os.path.join(default_dir, self.MAPPING_FILE_PATTERN))
if not default_files:
logger.warning(
f"EPGMapping: No default mapping files found in {default_dir}"
)
logger.warning(f"EPGMapping: No default mapping files found in {default_dir}")
return False
copied_count = 0
@@ -154,9 +138,7 @@ class EPGMapping:
logger.info(f"EPGMapping: Copied default mapping: {filename}")
copied_count += 1
else:
logger.error(
f"EPGMapping: Failed to write user mapping: {filename}"
)
logger.error(f"EPGMapping: Failed to write user mapping: {filename}")
except Exception as e:
logger.error(f"EPGMapping: Failed to read/copy {filename}: {e}")
@@ -215,18 +197,14 @@ class EPGMapping:
# Check if file exists
if not self.user_vfs.exists(filename):
logger.debug(
f"EPGMapping: No mapping file found for provider '{provider_name}'"
)
logger.debug(f"EPGMapping: No mapping file found for provider '{provider_name}'")
return None
try:
# Load mapping from file
raw_mapping = self.user_vfs.read_json(filename)
if raw_mapping is None:
logger.warning(
f"EPGMapping: Failed to parse mapping file for '{provider_name}'"
)
logger.warning(f"EPGMapping: Failed to parse mapping file for '{provider_name}'")
return None
# Flatten the mapping
@@ -248,9 +226,7 @@ class EPGMapping:
epg_id = channel_data
name = channel_id
else:
logger.warning(
f"EPGMapping: Invalid format for {provider_name}/{channel_id}"
)
logger.warning(f"EPGMapping: Invalid format for {provider_name}/{channel_id}")
continue
if epg_id:
@@ -266,16 +242,13 @@ class EPGMapping:
self._cache[provider_name] = cached_data
logger.info(
f"EPGMapping: Loaded mapping for '{provider_name}' "
f"with {len(mapping)} channels"
f"EPGMapping: Loaded mapping for '{provider_name}' " f"with {len(mapping)} channels"
)
return cached_data
except Exception as e:
logger.error(
f"EPGMapping: Failed to load mapping for '{provider_name}': {e}"
)
logger.error(f"EPGMapping: Failed to load mapping for '{provider_name}': {e}")
return None
def get_epg_channel_id(self, provider_name: str, channel_id: str) -> Optional[str]:
@@ -305,9 +278,7 @@ class EPGMapping:
epg_id = provider_data["mapping"].get(channel_id)
if epg_id:
logger.debug(
f"EPGMapping: Mapped '{provider_name}/{channel_id}' -> '{epg_id}'"
)
logger.debug(f"EPGMapping: Mapped '{provider_name}/{channel_id}' -> '{epg_id}'")
return epg_id
else:
# Channel not in mapping - fall back to channel_id
+14 -47
View File
@@ -50,9 +50,7 @@ class EPGParser:
provider_hash = int(provider_hash_obj.hexdigest()[:4], 16)
self._provider_registry[provider_hash] = provider_name
logger.debug(
f"Registered provider '{provider_name}' with hash {provider_hash:04x}"
)
logger.debug(f"Registered provider '{provider_name}' with hash {provider_hash:04x}")
return provider_hash
@@ -191,35 +189,23 @@ class EPGParser:
if len(parts) >= 2:
# Handle season number
season_str = parts[0].split("/")[0] if parts[0] else None
series_num = (
int(season_str) + 1
if season_str and season_str.strip()
else None
)
series_num = int(season_str) + 1 if season_str and season_str.strip() else None
# Handle episode number
episode_str = (
parts[1].split("/")[0] if len(parts) > 1 and parts[1] else None
)
episode_str = parts[1].split("/")[0] if len(parts) > 1 and parts[1] else None
episode_num = (
int(episode_str) + 1
if episode_str and episode_str.strip()
else None
int(episode_str) + 1 if episode_str and episode_str.strip() else None
)
# Handle part number (optional)
part_num = None
if len(parts) >= 3 and parts[2]:
part_str = parts[2].split("/")[0]
part_num = (
int(part_str) + 1 if part_str and part_str.strip() else None
)
part_num = int(part_str) + 1 if part_str and part_str.strip() else None
return series_num, episode_num, part_num
except (ValueError, IndexError) as e:
logger.debug(
f"Failed to parse xmltv_ns episode number '{episode_elem.text}': {e}"
)
logger.debug(f"Failed to parse xmltv_ns episode number '{episode_elem.text}': {e}")
elif system == "onscreen":
# Format: "S05E12" or "5x12"
@@ -304,11 +290,7 @@ class EPGParser:
return EPGGenre.MUSICBALLETDANCE
elif "arts" in genre_str or "culture" in genre_str or "art" in genre_str:
return EPGGenre.ARTSCULTURE
elif (
"educational" in genre_str
or "education" in genre_str
or "science" in genre_str
):
elif "educational" in genre_str or "education" in genre_str or "science" in genre_str:
return EPGGenre.EDUCATIONALSCIENCE
elif (
"social" in genre_str
@@ -362,14 +344,10 @@ class EPGParser:
# Get title
title_elem = programme_elem.find("title")
title = (
title_elem.text if title_elem is not None and title_elem.text else "Unknown"
)
title = title_elem.text if title_elem is not None and title_elem.text else "Unknown"
# Generate broadcast ID (with provider encoding if available)
broadcast_id = EPGParser.generate_broadcast_id(
epg_channel_id, start_time, provider_name
)
broadcast_id = EPGParser.generate_broadcast_id(epg_channel_id, start_time, provider_name)
# Build EPGEntry with required fields
entry_kwargs: Dict[str, Any] = {
@@ -452,9 +430,7 @@ class EPGParser:
if series_num is None:
for ep_elem in episode_nums:
if ep_elem.get("system") == "onscreen":
series_num, episode_num, part_num = EPGParser.parse_episode_num(
ep_elem
)
series_num, episode_num, part_num = EPGParser.parse_episode_num(ep_elem)
break
if series_num is not None:
@@ -555,9 +531,7 @@ class EPGParser:
if provider_hash not in self._provider_registry:
self._provider_registry[provider_hash] = provider_name
logger.debug(
f"Registered provider '{provider_name}' with hash {provider_hash:04x}"
)
logger.debug(f"Registered provider '{provider_name}' with hash {provider_hash:04x}")
programmes: List[EPGEntry] = []
@@ -571,16 +545,11 @@ class EPGParser:
# Check if this programme matches our channel
if elem.get("channel") == epg_channel_id:
# Parse the programme WITHOUT calling register_provider again
epg_entry = self.parse_programme(
elem, epg_channel_id, provider_name
)
epg_entry = self.parse_programme(elem, epg_channel_id, provider_name)
if epg_entry:
# Filter by time range using EPGEntry methods
if (
start_time is not None
and epg_entry.end < start_time
):
if start_time is not None and epg_entry.end < start_time:
# Programme ends before requested range
elem.clear()
continue
@@ -596,9 +565,7 @@ class EPGParser:
# Clear element to free memory
elem.clear()
logger.info(
f"Parsed {len(programmes)} programmes for channel '{epg_channel_id}'"
)
logger.info(f"Parsed {len(programmes)} programmes for channel '{epg_channel_id}'")
return programmes
except ET.ParseError as e:
@@ -18,9 +18,7 @@ class EPGOperations:
self.epg_manager = EPGManager()
logger.debug("EPGOperations: Initialized")
def get_channel_epg(
self, provider_name: str, channel_id: str, **kwargs
) -> List[Dict]:
def get_channel_epg(self, provider_name: str, channel_id: str, **kwargs) -> List[Dict]:
"""Get EPG data for a specific channel."""
provider = self.registry.get_provider(provider_name)
if not provider:
+10 -30
View File
@@ -50,9 +50,7 @@ class ProviderManager:
def discover_all_providers(self, default_country: str = "DE") -> List[str]:
return self.registry.discover_all_providers(default_country)
def discover_providers(
self, country: str = "DE", detected_providers: Dict = None
) -> List[str]:
def discover_providers(self, country: str = "DE", detected_providers: Dict = None) -> List[str]:
"""Legacy method for backward compatibility."""
if not self.registry.provider_metadata:
self.registry.discover_all_providers(country)
@@ -71,9 +69,7 @@ class ProviderManager:
return self.registry.reinitialize_provider(provider_name)
def reinitialize_providers(self, provider_names: List[str]) -> Dict[str, bool]:
return {
name: self.registry.reinitialize_provider(name) for name in provider_names
}
return {name: self.registry.reinitialize_provider(name) for name in provider_names}
def reinitialize_all_providers(self) -> Dict[str, bool]:
enabled = self.registry.get_enabled_providers()
@@ -111,12 +107,8 @@ class ProviderManager:
) -> List[StreamingChannel]:
return self.channel_ops.get_channels(provider_name, fetch_manifests, **kwargs)
def get_channel_manifest(
self, provider_name: str, channel_id: str, **kwargs
) -> Optional[str]:
return self.channel_ops.get_channel_manifest(
provider_name, channel_id, **kwargs
)
def get_channel_manifest(self, provider_name: str, channel_id: str, **kwargs) -> Optional[str]:
return self.channel_ops.get_channel_manifest(provider_name, channel_id, **kwargs)
def get_all_channels(
self, fetch_manifests: bool = True, **kwargs
@@ -127,9 +119,7 @@ class ProviderManager:
# EPG OPERATIONS (delegate to EPGOperations)
# ==========================================================================
def get_channel_epg(
self, provider_name: str, channel_id: str, **kwargs
) -> List[Dict]:
def get_channel_epg(self, provider_name: str, channel_id: str, **kwargs) -> List[Dict]:
return self.epg_ops.get_channel_epg(provider_name, channel_id, **kwargs)
def get_provider_epg_xmltv(self, provider_name: str, **kwargs) -> Optional[str]:
@@ -154,9 +144,7 @@ class ProviderManager:
# DRM OPERATIONS (delegate to DRMOperations)
# ==========================================================================
def get_channel_drm_configs(
self, provider_name: str, channel_id: str, **kwargs
) -> List:
def get_channel_drm_configs(self, provider_name: str, channel_id: str, **kwargs) -> List:
return self.drm_ops.get_channel_drm_configs(provider_name, channel_id, **kwargs)
def list_drm_plugins(self) -> Dict:
@@ -195,9 +183,7 @@ class ProviderManager:
provider_name, channel_id, start_time, end_time, epg_id, country
)
def get_catchup_window(
self, provider_name: str, channel_id: Optional[str] = None
) -> int:
def get_catchup_window(self, provider_name: str, channel_id: Optional[str] = None) -> int:
return self.catchup_ops.get_catchup_window(provider_name, channel_id)
def supports_catchup(self, provider_name: str) -> bool:
@@ -220,9 +206,7 @@ class ProviderManager:
return self.subscription_ops.get_available_packages(provider_name, **kwargs)
def is_channel_accessible(self, provider_name: str, channel_id: str, **kwargs):
return self.subscription_ops.is_channel_accessible(
provider_name, channel_id, **kwargs
)
return self.subscription_ops.is_channel_accessible(provider_name, channel_id, **kwargs)
# ==========================================================================
# UTILITY METHODS (remain in ProviderManager as helpers)
@@ -237,9 +221,7 @@ class ProviderManager:
http_manager = provider.http_manager
if not http_manager:
logger.warning(
f"ProviderManager: Provider '{provider_name}' has no HTTP manager"
)
logger.warning(f"ProviderManager: Provider '{provider_name}' has no HTTP manager")
return None
logger.debug(f"ProviderManager: Retrieved HTTP manager for '{provider_name}'")
@@ -280,9 +262,7 @@ class ProviderManager:
available = self.registry.get_enabled_providers()
if not available:
logger.warning(
"ProviderManager: No enabled providers available for selection"
)
logger.warning("ProviderManager: No enabled providers available for selection")
return []
selected = []
@@ -1,6 +1,5 @@
# streaming_providers/base/models/__init__.py
from .drm_models import (DRMConfig, DRMSystem, LicenseConfig,
LicenseUnwrapperParams)
from .drm_models import DRMConfig, DRMSystem, LicenseConfig, LicenseUnwrapperParams
from .streaming_channel import StreamingChannel
from .subscription import SubscriptionPackage, UserSubscription
+2 -6
View File
@@ -57,9 +57,7 @@ class AuthStatus:
result = {
"provider": (
f"{self.provider_name}_{self.country}"
if self.country
else self.provider_name
f"{self.provider_name}_{self.country}" if self.country else self.provider_name
),
"provider_name": self.provider_name,
"provider_label": self.provider_label,
@@ -101,9 +99,7 @@ class AuthStatus:
result["refresh_token_expires_at"] = self.refresh_token_expires_at
if self.refresh_token_expires_in_seconds is not None:
result["refresh_token_expires_in_seconds"] = (
self.refresh_token_expires_in_seconds
)
result["refresh_token_expires_in_seconds"] = self.refresh_token_expires_in_seconds
return result
@@ -45,23 +45,17 @@ class WrapperType(str, Enum):
NONE = "none"
class UnwrapperType(str, Enum):
AUTO = "auto"
BASE64 = "base64"
JSON = "json"
XML = "xml"
NONE = "none"
@dataclass
class PSSHData:
"""
Protection System Specific Header data for DRM systems
"""
system_id: str
pssh_box: str = "" # Base64 encoded PSSH box
key_ids: List[str] = field(default_factory=list)
source: str = "manifest" # "manifest", "mp4_segment", "unknown"
system_id: str # UUID of the DRM system
pssh_box: str # Base64 encoded PSSH box data
key_ids: List[str] = field(default_factory=list) # Optional key IDs
@property
def needs_extraction(self) -> bool:
"""Check if PSSH/key_ids need to be extracted from segments"""
return not self.pssh_box or not self.key_ids
@property
def drm_system(self) -> Optional[DRMSystem]:
@@ -73,13 +67,12 @@ class PSSHData:
if not self.system_id:
raise ValueError("system_id is required")
if not self.pssh_box:
raise ValueError("pssh_box is required")
try:
base64.b64decode(self.pssh_box)
except Exception:
raise ValueError("pssh_box must be valid base64")
# FIX: pssh_box can be empty (PSSH in segments)
if self.pssh_box: # Only validate if not empty
try:
base64.b64decode(self.pssh_box)
except Exception:
raise ValueError("pssh_box must be valid base64")
# Validate key IDs if present
for kid in self.key_ids:
@@ -87,6 +80,14 @@ class PSSHData:
raise ValueError(f"Invalid key ID format: {kid}")
class UnwrapperType(str, Enum):
AUTO = "auto"
BASE64 = "base64"
JSON = "json"
XML = "xml"
NONE = "none"
@dataclass
class LicenseUnwrapperParams:
path_data: Optional[str] = None
@@ -136,9 +137,7 @@ class LicenseConfig:
"""Helper to ensure req_data is base64 encoded"""
import base64
req_data_encoded = base64.b64encode(req_data_template.encode("utf-8")).decode(
"utf-8"
)
req_data_encoded = base64.b64encode(req_data_template.encode("utf-8")).decode("utf-8")
return cls(req_data=req_data_encoded, **kwargs)
@@ -181,9 +180,7 @@ class DRMConfig:
license_dict["unwrapper"] = self.license.unwrapper
if self.license.unwrapper_params:
license_dict["unwrapper_params"] = {
k: v
for k, v in vars(self.license.unwrapper_params).items()
if v is not None
k: v for k, v in vars(self.license.unwrapper_params).items() if v is not None
}
if self.license.keyids:
license_dict["keyids"] = self.license.keyids
@@ -400,16 +400,10 @@ class EPGEntry:
"""
if not text:
return []
return [
item.strip()
for item in text.split(EPG_STRING_TOKEN_SEPARATOR)
if item.strip()
]
return [item.strip() for item in text.split(EPG_STRING_TOKEN_SEPARATOR) if item.strip()]
@staticmethod
def encode_broadcast_id(
provider_name: str, channel_id: str, start_time: int
) -> int:
def encode_broadcast_id(provider_name: str, channel_id: str, start_time: int) -> int:
"""
Generate deterministic broadcast ID with encoded provider information.
@@ -526,18 +520,12 @@ class EPGEntry:
raise ValueError("end time must be after start time")
# Validate episode numbers if set
if (
self.season_number is not None
and self.season_number < EPG_TAG_INVALID_SERIES_EPISODE
):
if self.season_number is not None and self.season_number < EPG_TAG_INVALID_SERIES_EPISODE:
raise ValueError(
f"season_number must be >= EPG_TAG_INVALID_SERIES_EPISODE ({EPG_TAG_INVALID_SERIES_EPISODE})"
)
if (
self.episode_number is not None
and self.episode_number < EPG_TAG_INVALID_SERIES_EPISODE
):
if self.episode_number is not None and self.episode_number < EPG_TAG_INVALID_SERIES_EPISODE:
raise ValueError(
f"episode_number must be >= EPG_TAG_INVALID_SERIES_EPISODE ({EPG_TAG_INVALID_SERIES_EPISODE})"
)
@@ -144,9 +144,7 @@ class ProxyConfig:
"""Create ProxyConfig from dictionary"""
auth = None
if "auth" in data and data["auth"]:
auth = ProxyAuth(
username=data["auth"]["username"], password=data["auth"]["password"]
)
auth = ProxyAuth(username=data["auth"]["username"], password=data["auth"]["password"])
scope_data = data.get("scope", {})
scope = ProxyScope(
@@ -169,9 +167,7 @@ class ProxyConfig:
)
@classmethod
def from_url(
cls, proxy_url: str, scope: Optional[ProxyScope] = None
) -> "ProxyConfig":
def from_url(cls, proxy_url: str, scope: Optional[ProxyScope] = None) -> "ProxyConfig":
"""
Create ProxyConfig from proxy URL string
@@ -252,9 +248,7 @@ class RequestConfig:
}
# Add proxy if configured and enabled for this operation
if self.proxy_config and self.proxy_config.scope.should_use_proxy_for(
operation
):
if self.proxy_config and self.proxy_config.scope.should_use_proxy_for(operation):
kwargs["proxies"] = self.proxy_config.to_proxy_dict()
return kwargs
@@ -282,9 +282,7 @@ class StreamingChannel:
self.is_radio = True
# Update quality if not set
if not self.quality or self.quality.upper() not in [
q.value for q in Quality
]:
if not self.quality or self.quality.upper() not in [q.value for q in Quality]:
self.quality = "AUDIO"
# Update content_type if it's still LIVE
@@ -122,9 +122,7 @@ class HTTPManager:
"""Perform DELETE request with proxy support"""
return self._make_request("DELETE", url, operation, **kwargs)
def _make_request(
self, method: str, url: str, operation: str, **kwargs
) -> requests.Response:
def _make_request(self, method: str, url: str, operation: str, **kwargs) -> requests.Response:
"""
Make HTTP request with full configuration support
"""
@@ -182,9 +180,7 @@ class HTTPManager:
)
raise
def _log_request(
self, method: str, url: str, operation: str, kwargs: Dict[str, Any]
) -> None:
def _log_request(self, method: str, url: str, operation: str, kwargs: Dict[str, Any]) -> None:
"""Log request details with comprehensive proxy information"""
# Build proxy information string
@@ -192,13 +188,9 @@ class HTTPManager:
if self.config.proxy_config:
if self.config.proxy_config.scope.should_use_proxy_for(operation):
# Proxy is configured and will be used
proxy_host = (
f"{self.config.proxy_config.host}:{self.config.proxy_config.port}"
)
proxy_host = f"{self.config.proxy_config.host}:{self.config.proxy_config.port}"
proxy_type = self.config.proxy_config.proxy_type.value
has_auth = (
"authenticated" if self.config.proxy_config.auth else "no-auth"
)
has_auth = "authenticated" if self.config.proxy_config.auth else "no-auth"
proxy_info = f" [proxy: {proxy_type}://{proxy_host} ({has_auth})]"
else:
# Proxy is configured but not used for this operation
@@ -348,9 +340,7 @@ class HTTPManagerFactory:
return HTTPManager(config)
@staticmethod
def create_with_proxy_url(
provider_name: str, proxy_url: str, **kwargs
) -> HTTPManager:
def create_with_proxy_url(provider_name: str, proxy_url: str, **kwargs) -> HTTPManager:
"""
Create HTTP manager with proxy from URL string
@@ -363,6 +353,4 @@ class HTTPManagerFactory:
Configured HTTPManager instance
"""
proxy_config = ProxyConfig.from_url(proxy_url)
return HTTPManagerFactory.create_for_provider(
provider_name, proxy_config, **kwargs
)
return HTTPManagerFactory.create_for_provider(provider_name, proxy_config, **kwargs)
@@ -62,9 +62,7 @@ class ProxyConfigManager:
vfs_data = self.vfs.read_json(self.proxy_config_file)
if vfs_data is None:
logger.debug(
"No proxy configuration file found, starting with empty config"
)
logger.debug("No proxy configuration file found, starting with empty config")
return
# Load global configuration
@@ -78,27 +76,18 @@ class ProxyConfigManager:
try:
# Check if this is a nested (country-aware) structure
if isinstance(provider_data, dict) and any(
isinstance(v, dict) and len(k) <= 3
for k, v in provider_data.items()
isinstance(v, dict) and len(k) <= 3 for k, v in provider_data.items()
):
# Country-aware structure
for country, config_data in provider_data.items():
if isinstance(config_data, dict) and len(country) <= 3:
cache_key = f"{provider_name}_{country}"
self._config_cache[cache_key] = ProxyConfig.from_dict(
config_data
)
logger.debug(
f"Loaded proxy config for {provider_name} ({country})"
)
self._config_cache[cache_key] = ProxyConfig.from_dict(config_data)
logger.debug(f"Loaded proxy config for {provider_name} ({country})")
else:
# Flat structure (no country)
self._config_cache[provider_name] = ProxyConfig.from_dict(
provider_data
)
logger.debug(
f"Loaded proxy configuration for provider: {provider_name}"
)
self._config_cache[provider_name] = ProxyConfig.from_dict(provider_data)
logger.debug(f"Loaded proxy configuration for provider: {provider_name}")
except Exception as e:
logger.error(f"Error loading proxy config for {provider_name}: {e}")
@@ -230,22 +219,16 @@ class ProxyConfigManager:
cache_key, _ = self._get_proxy_path(provider_name, country)
self._config_cache[cache_key] = proxy_config
country_str = f" ({country})" if country else ""
logger.info(
f"Set proxy configuration for provider: {provider_name}{country_str}"
)
logger.info(f"Set proxy configuration for provider: {provider_name}{country_str}")
return self._save_configurations()
except Exception as e:
country_str = f" ({country})" if country else ""
logger.error(
f"Error setting proxy configuration for {provider_name}{country_str}: {e}"
)
logger.error(f"Error setting proxy configuration for {provider_name}{country_str}: {e}")
return False
def remove_proxy_config(
self, provider_name: str, country: Optional[str] = None
) -> bool:
def remove_proxy_config(self, provider_name: str, country: Optional[str] = None) -> bool:
"""
Remove proxy configuration for a provider
@@ -266,9 +249,7 @@ class ProxyConfigManager:
cache_key, _ = self._get_proxy_path(provider_name, country)
if cache_key in self._config_cache:
del self._config_cache[cache_key]
logger.info(
f"Removed proxy configuration for {provider_name} ({country})"
)
logger.info(f"Removed proxy configuration for {provider_name} ({country})")
else:
logger.warning(
f"No proxy configuration found for {provider_name} ({country})"
@@ -285,13 +266,9 @@ class ProxyConfigManager:
if keys_to_remove:
for key in keys_to_remove:
del self._config_cache[key]
logger.info(
f"Removed all proxy configurations for {provider_name}"
)
logger.info(f"Removed all proxy configurations for {provider_name}")
else:
logger.warning(
f"No proxy configuration found for {provider_name}"
)
logger.warning(f"No proxy configuration found for {provider_name}")
return True
return self._save_configurations()
@@ -372,9 +349,7 @@ class ProxyConfigManager:
manager.close()
return result
def get_proxy_info(
self, provider_name: str, country: Optional[str] = None
) -> Dict[str, Any]:
def get_proxy_info(self, provider_name: str, country: Optional[str] = None) -> Dict[str, Any]:
"""
Get detailed information about proxy configuration
@@ -431,9 +406,7 @@ class ProxyConfigManager:
Path to exported file
"""
if not export_path:
export_path = str(
self.config_dir / f"proxy_config_backup_{int(time.time())}.json"
)
export_path = str(self.config_dir / f"proxy_config_backup_{int(time.time())}.json")
try:
# Create export data with country-aware structure
@@ -459,9 +432,7 @@ class ProxyConfigManager:
"version": "1.1",
"source": "streaming_providers_proxy_manager",
},
"global": (
self._global_config.to_dict() if self._global_config else None
),
"global": (self._global_config.to_dict() if self._global_config else None),
"providers": providers_export,
}
@@ -505,30 +476,23 @@ class ProxyConfigManager:
try:
# Check if nested (country-aware)
if isinstance(provider_data, dict) and any(
isinstance(v, dict) and len(k) <= 3
for k, v in provider_data.items()
isinstance(v, dict) and len(k) <= 3 for k, v in provider_data.items()
):
# Country-aware structure
for country, config_data in provider_data.items():
if isinstance(config_data, dict) and len(country) <= 3:
cache_key = f"{provider_name}_{country}"
self._config_cache[cache_key] = ProxyConfig.from_dict(
config_data
)
self._config_cache[cache_key] = ProxyConfig.from_dict(config_data)
else:
# Flat structure
self._config_cache[provider_name] = ProxyConfig.from_dict(
provider_data
)
self._config_cache[provider_name] = ProxyConfig.from_dict(provider_data)
except Exception as e:
logger.error(f"Error importing config for {provider_name}: {e}")
# Save the imported configurations
success = self._save_configurations()
if success:
logger.info(
f"Successfully imported proxy configurations from {import_path}"
)
logger.info(f"Successfully imported proxy configurations from {import_path}")
return success
except Exception as e:
@@ -583,9 +547,7 @@ class ProxyConfigManager:
"""
results = {}
for provider_name in provider_names:
results[provider_name] = self.set_proxy_config(
provider_name, proxy_config, country
)
results[provider_name] = self.set_proxy_config(provider_name, proxy_config, country)
return results
def get_all_proxy_info(self) -> Dict[str, Dict[str, Any]]:
+19 -49
View File
@@ -264,9 +264,7 @@ class StreamingProvider(ABC):
"""Get manifest URL for a specific channel by ID"""
return None
def get_dynamic_manifest_params(
self, channel: StreamingChannel, **kwargs
) -> Optional[str]:
def get_dynamic_manifest_params(self, channel: StreamingChannel, **kwargs) -> Optional[str]:
"""Optional: Get dynamic manifest parameters for a channel"""
return None
@@ -283,9 +281,7 @@ class StreamingProvider(ABC):
def to_json(self, channels: List[StreamingChannel] = None, indent: int = 2) -> str:
"""Convert to JSON string"""
return json.dumps(
self.to_output_format(channels), indent=indent, ensure_ascii=False
)
return json.dumps(self.to_output_format(channels), indent=indent, ensure_ascii=False)
# ============================================================================
# HTTP MANAGER SETUP (Already Implemented)
@@ -351,9 +347,7 @@ class StreamingProvider(ABC):
) -> Optional[ProxyConfig]:
"""Resolve proxy configuration from multiple sources with priority"""
if proxy_config is not None:
logger.debug(
f"{provider_name}: Using directly provided proxy configuration"
)
logger.debug(f"{provider_name}: Using directly provided proxy configuration")
return proxy_config
if proxy_url:
@@ -361,9 +355,7 @@ class StreamingProvider(ABC):
logger.debug(f"{provider_name}: Creating proxy config from URL")
return ProxyConfig.from_url(proxy_url)
except Exception as e:
logger.warning(
f"{provider_name}: Failed to parse proxy URL '{proxy_url}': {e}"
)
logger.warning(f"{provider_name}: Failed to parse proxy URL '{proxy_url}': {e}")
try:
from .network import ProxyConfigManager
@@ -375,14 +367,10 @@ class StreamingProvider(ABC):
logger.debug(f"{provider_name}: Using proxy from ProxyConfigManager")
return managed_proxy
else:
logger.debug(
f"{provider_name}: No proxy configuration found in ProxyConfigManager"
)
logger.debug(f"{provider_name}: No proxy configuration found in ProxyConfigManager")
except Exception as e:
logger.warning(
f"{provider_name}: Could not load proxy from ProxyConfigManager: {e}"
)
logger.warning(f"{provider_name}: Could not load proxy from ProxyConfigManager: {e}")
logger.debug(f"{provider_name}: No proxy configuration available")
return None
@@ -395,9 +383,7 @@ class StreamingProvider(ABC):
info_parts = [f"HTTP manager initialized for '{provider_name}'"]
if proxy_config:
proxy_type = (
proxy_config.proxy_type.value if proxy_config.proxy_type else "http"
)
proxy_type = proxy_config.proxy_type.value if proxy_config.proxy_type else "http"
proxy_host = f"{proxy_config.host}:{proxy_config.port}"
has_auth = "authenticated" if proxy_config.auth else "no-auth"
info_parts.append(f"proxy: {proxy_type}://{proxy_host} ({has_auth})")
@@ -429,14 +415,10 @@ class StreamingProvider(ABC):
if http_manager and hasattr(authenticator, "http_manager"):
if authenticator.http_manager is None:
logger.debug(
f"{self.provider_name}: Sharing HTTP manager with authenticator"
)
logger.debug(f"{self.provider_name}: Sharing HTTP manager with authenticator")
authenticator.http_manager = http_manager
else:
logger.debug(
f"{self.provider_name}: Using authenticator's existing HTTP manager"
)
logger.debug(f"{self.provider_name}: Using authenticator's existing HTTP manager")
http_manager = authenticator.http_manager
return http_manager
@@ -512,9 +494,7 @@ class StreamingProvider(ABC):
elif self.authenticator is not None:
token = self.authenticator.get_bearer_token(**kwargs)
else:
logger.warning(
f"{self.provider_name}: No token getter or authenticator available"
)
logger.warning(f"{self.provider_name}: No token getter or authenticator available")
token = None
# Add auth header based on type
@@ -643,17 +623,13 @@ class StreamingProvider(ABC):
try:
# Try to get token based on type
if token_type == "bearer":
return self.authenticator.get_bearer_token(
force_refresh=force_refresh, **kwargs
)
return self.authenticator.get_bearer_token(force_refresh=force_refresh, **kwargs)
elif hasattr(self.authenticator, f"get_{token_type}_token"):
getter = getattr(self.authenticator, f"get_{token_type}_token")
return getter(force_refresh=force_refresh, **kwargs)
else:
# Default to bearer token
return self.authenticator.get_bearer_token(
force_refresh=force_refresh, **kwargs
)
return self.authenticator.get_bearer_token(force_refresh=force_refresh, **kwargs)
except Exception as e:
logger.error(f"{self.provider_name}: Error getting {token_type} token: {e}")
return None
@@ -723,9 +699,7 @@ class StreamingProvider(ABC):
Override in subclass if catchup requires different DRM configuration.
"""
if not self.supports_catchup:
logger.debug(
f"{self.provider_name}: Catchup not supported, falling back to live DRM"
)
logger.debug(f"{self.provider_name}: Catchup not supported, falling back to live DRM")
return self.get_drm(channel_id, **kwargs)
logger.debug(
@@ -1060,16 +1034,13 @@ class StreamingProvider(ABC):
ValueError: If auth_type is not supported
"""
if not self.validate_auth_type(auth_type):
raise ValueError(
f"Auth type '{auth_type}' not supported by {self.provider_name}"
)
raise ValueError(f"Auth type '{auth_type}' not supported by {self.provider_name}")
requirements = {
"auth_type": auth_type,
"needs_storage": auth_type in ["user_credentials", "client_credentials"],
"provides_token": auth_type != "anonymous",
"user_interaction_required": auth_type
in ["user_credentials", "device_registration"],
"user_interaction_required": auth_type in ["user_credentials", "device_registration"],
}
# Type-specific details
@@ -1116,9 +1087,7 @@ class StreamingProvider(ABC):
def requires_stored_credentials(self) -> bool:
"""True if provider needs credentials stored in settings."""
credential_types = ["user_credentials", "client_credentials"]
return any(
auth_type in credential_types for auth_type in self.supported_auth_types
)
return any(auth_type in credential_types for auth_type in self.supported_auth_types)
# ===== AUTHENTICATION PROPERTIES =====
@@ -1207,8 +1176,9 @@ class StreamingProvider(ABC):
Returns:
AuthStatus object
"""
from ..providers.auth_builder import \
AuthStatusBuilder # Import here to avoid circular imports
from ..providers.auth_builder import (
AuthStatusBuilder,
) # Import here to avoid circular imports
return AuthStatusBuilder.for_provider(self, context)
@@ -161,9 +161,7 @@ class ProviderRegistry:
if instance:
self.providers[instance_name] = instance
logger.info(
f"ProviderRegistry: Discovered {len(discovered)} provider instances"
)
logger.info(f"ProviderRegistry: Discovered {len(discovered)} provider instances")
return discovered
def get_provider(self, provider_name: str) -> Optional[StreamingProvider]:
@@ -199,9 +197,7 @@ class ProviderRegistry:
from .settings.provider_enable_manager import ProviderEnableManager
enable_manager = ProviderEnableManager()
success, message = enable_manager.set_provider_enabled(
provider_name, enabled
)
success, message = enable_manager.set_provider_enabled(provider_name, enabled)
return success
except Exception as e:
logger.error(f"Error updating enable status: {e}")
@@ -1,7 +1,6 @@
# streaming_providers/base/settings/__init__.py
from .kodi_settings_bridge import KodiSettingsBridge
from .models.provider_settings import (ProviderSettingsSchema,
StandardProviderSettings)
from .models.provider_settings import ProviderSettingsSchema, StandardProviderSettings
from .models.settings_models import SettingType, SettingValue, ValidationRule
from .settings_manager import UnifiedSettingsManager
@@ -3,11 +3,9 @@ import json
import xml.etree.ElementTree as ElementTree
from typing import Any, Dict, List, Optional, Set, Tuple
from ..auth.credentials import (BaseCredentials, ClientCredentials,
UserPasswordCredentials)
from ..auth.credentials import BaseCredentials, ClientCredentials, UserPasswordCredentials
from ..models.proxy_models import ProxyConfig
from ..utils.environment import (get_environment_manager, get_vfs_instance,
is_kodi_environment)
from ..utils.environment import get_environment_manager, get_vfs_instance, is_kodi_environment
from ..utils.logger import logger
@@ -17,9 +15,7 @@ class KodiSettingsBridge:
# Markers that identify credential settings
CREDENTIAL_MARKERS = {"_username", "_password", "_client_id", "_client_secret"}
def __init__(
self, addon_id: Optional[str] = None, config_dir: Optional[str] = None
):
def __init__(self, addon_id: Optional[str] = None, config_dir: Optional[str] = None):
"""Initialize Kodi settings bridge"""
self.addon = None
self.addon_id = addon_id
@@ -42,9 +38,7 @@ class KodiSettingsBridge:
else:
self.addon = xbmcaddon.Addon()
self.addon_id = self.addon.getAddonInfo("id")
logger.info(
f"Kodi settings bridge initialized for addon: {self.addon_id}"
)
logger.info(f"Kodi settings bridge initialized for addon: {self.addon_id}")
except Exception as e:
logger.error(f"Failed to initialize Kodi addon: {e}")
self.addon = None
@@ -60,9 +54,7 @@ class KodiSettingsBridge:
content = self.vfs.read_text(self._settings_file)
if content:
self._standalone_settings = json.loads(content)
logger.debug(
f"Loaded {len(self._standalone_settings)} standalone settings"
)
logger.debug(f"Loaded {len(self._standalone_settings)} standalone settings")
except Exception as e:
logger.error(f"Error loading standalone settings: {e}")
@@ -71,9 +63,7 @@ class KodiSettingsBridge:
if not is_kodi_environment():
try:
self.vfs.write_json(self._settings_file, self._standalone_settings)
logger.debug(
f"Saved {len(self._standalone_settings)} standalone settings"
)
logger.debug(f"Saved {len(self._standalone_settings)} standalone settings")
except Exception as e:
logger.error(f"Error saving standalone settings: {e}")
@@ -166,9 +156,7 @@ class KodiSettingsBridge:
logger.error(f"Error reading settings.xml: {e}")
return []
def _parse_provider_country(
self, setting_id: str
) -> Optional[Tuple[str, Optional[str]]]:
def _parse_provider_country(self, setting_id: str) -> Optional[Tuple[str, Optional[str]]]:
"""
Parse a setting ID to extract provider and optional country.
@@ -312,9 +300,7 @@ class KodiSettingsBridge:
username = self.get_setting(f"{provider}{country_suffix}_username")
password = self.get_setting(f"{provider}{country_suffix}_password")
client_id = self.get_setting(f"{provider}{country_suffix}_client_id")
client_secret = self.get_setting(
f"{provider}{country_suffix}_client_secret"
)
client_secret = self.get_setting(f"{provider}{country_suffix}_client_secret")
logger.debug(f"Settings for {provider}{country_suffix}:")
logger.debug(f" username: '{username}' (empty={not username})")
@@ -324,9 +310,7 @@ class KodiSettingsBridge:
# Determine credential type based on available values
if username and password:
logger.info(
f"Found username/password credentials for {provider}{country_suffix}"
)
logger.info(f"Found username/password credentials for {provider}{country_suffix}")
return UserPasswordCredentials(
username=username.strip(),
password=password.strip(),
@@ -353,25 +337,15 @@ class KodiSettingsBridge:
try:
if isinstance(credentials, UserPasswordCredentials):
self.set_setting(
f"{provider}{country_suffix}_username", credentials.username
)
self.set_setting(
f"{provider}{country_suffix}_password", credentials.password
)
self.set_setting(f"{provider}{country_suffix}_username", credentials.username)
self.set_setting(f"{provider}{country_suffix}_password", credentials.password)
if credentials.client_id:
self.set_setting(
f"{provider}{country_suffix}_client_id", credentials.client_id
)
logger.info(
f"Wrote username/password credentials for {provider}{country_suffix}"
)
self.set_setting(f"{provider}{country_suffix}_client_id", credentials.client_id)
logger.info(f"Wrote username/password credentials for {provider}{country_suffix}")
return True
elif isinstance(credentials, ClientCredentials):
self.set_setting(
f"{provider}{country_suffix}_client_id", credentials.client_id
)
self.set_setting(f"{provider}{country_suffix}_client_id", credentials.client_id)
self.set_setting(
f"{provider}{country_suffix}_client_secret",
credentials.client_secret,
@@ -418,12 +392,8 @@ class KodiSettingsBridge:
try:
# Check if proxy is enabled
proxy_enabled = self.get_setting(
f"{provider}{country_suffix}_proxy_enabled"
)
logger.debug(
f"Proxy enabled setting for {provider}{country_suffix}: '{proxy_enabled}'"
)
proxy_enabled = self.get_setting(f"{provider}{country_suffix}_proxy_enabled")
logger.debug(f"Proxy enabled setting for {provider}{country_suffix}: '{proxy_enabled}'")
if not proxy_enabled or proxy_enabled.lower() not in ["true", "1", "yes"]:
logger.debug(f"Proxy not enabled for {provider}{country_suffix}")
@@ -437,9 +407,7 @@ class KodiSettingsBridge:
logger.debug(f" port: '{proxy_port_str}'")
if not proxy_host or not proxy_port_str:
logger.debug(
f"Proxy host or port missing for {provider}{country_suffix}"
)
logger.debug(f"Proxy host or port missing for {provider}{country_suffix}")
return None
try:
@@ -470,12 +438,8 @@ class KodiSettingsBridge:
try:
self.set_setting(f"{provider}{country_suffix}_proxy_enabled", "true")
self.set_setting(
f"{provider}{country_suffix}_proxy_host", proxy_config.host
)
self.set_setting(
f"{provider}{country_suffix}_proxy_port", str(proxy_config.port)
)
self.set_setting(f"{provider}{country_suffix}_proxy_host", proxy_config.host)
self.set_setting(f"{provider}{country_suffix}_proxy_port", str(proxy_config.port))
logger.info(f"Wrote proxy config for {provider}{country_suffix}")
return True
@@ -523,9 +487,7 @@ class KodiSettingsBridge:
try:
ip_address = self.get_setting(f"{provider}{country_suffix}_ipaddress")
logger.debug(
f"IP address setting for {provider}{country_suffix}: '{ip_address}'"
)
logger.debug(f"IP address setting for {provider}{country_suffix}: '{ip_address}'")
if ip_address and ip_address.strip():
logger.info(
@@ -562,20 +524,13 @@ class KodiSettingsBridge:
and cred1.password == cred2.password
and cred1.client_id == cred2.client_id
)
elif isinstance(cred1, ClientCredentials) and isinstance(
cred2, ClientCredentials
):
return (
cred1.client_id == cred2.client_id
and cred1.client_secret == cred2.client_secret
)
elif isinstance(cred1, ClientCredentials) and isinstance(cred2, ClientCredentials):
return cred1.client_id == cred2.client_id and cred1.client_secret == cred2.client_secret
return False
@staticmethod
def _proxy_configs_equal(
proxy1: Optional[ProxyConfig], proxy2: Optional[ProxyConfig]
) -> bool:
def _proxy_configs_equal(proxy1: Optional[ProxyConfig], proxy2: Optional[ProxyConfig]) -> bool:
"""Compare two proxy configurations for equality"""
if proxy1 is None and proxy2 is None:
return True
@@ -1,10 +1,20 @@
# streaming_providers/base/settings/models/__init__.py
from .provider_settings import ProviderSettingsSchema, StandardProviderSettings
from .settings_models import (SettingType, SettingValue, SettingValueBuilder,
StandardValidationRules, ValidationRule,
boolean_setting, integer_setting, ip_setting,
password_setting, port_setting, select_setting,
string_setting, url_setting)
from .settings_models import (
SettingType,
SettingValue,
SettingValueBuilder,
StandardValidationRules,
ValidationRule,
boolean_setting,
integer_setting,
ip_setting,
password_setting,
port_setting,
select_setting,
string_setting,
url_setting,
)
__all__ = [
# From settings_models.py
@@ -3,9 +3,15 @@ from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Set
from ...utils.logger import logger
from .settings_models import (SettingValue, boolean_setting, integer_setting,
password_setting, port_setting, select_setting,
string_setting)
from .settings_models import (
SettingValue,
boolean_setting,
integer_setting,
password_setting,
port_setting,
select_setting,
string_setting,
)
@dataclass
@@ -361,9 +367,7 @@ class ProviderSettingsSchema:
required = []
for setting_name, setting in self._settings.items():
# Check if setting has a "not_empty" validation rule
has_required_rule = any(
rule.name == "not_empty" for rule in setting.validation_rules
)
has_required_rule = any(rule.name == "not_empty" for rule in setting.validation_rules)
if has_required_rule:
required.append(setting_name)
return required
@@ -468,12 +472,9 @@ class ProviderSettingsSchema:
"""Convert schema to dictionary representation"""
return {
"provider_name": self.provider_name,
"settings": {
name: setting.to_dict() for name, setting in self._settings.items()
},
"settings": {name: setting.to_dict() for name, setting in self._settings.items()},
"categories": {
category: list(settings)
for category, settings in self._categories.items()
category: list(settings) for category, settings in self._categories.items()
},
"kodi_mapping": self._kodi_mapping.copy(),
"configuration_status": self.get_configuration_completeness(),
@@ -576,17 +577,13 @@ class StandardProviderSettings:
return cls._registered_schemas["ard"]
@classmethod
def register_provider_schema(
cls, provider_name: str, schema: ProviderSettingsSchema
) -> None:
def register_provider_schema(cls, provider_name: str, schema: ProviderSettingsSchema) -> None:
"""Register a custom provider schema"""
cls._registered_schemas[provider_name] = schema
logger.info(f"Registered custom settings schema for provider: {provider_name}")
@classmethod
def get_provider_schema(
cls, provider_name: str
) -> Optional[ProviderSettingsSchema]:
def get_provider_schema(cls, provider_name: str) -> Optional[ProviderSettingsSchema]:
"""Get schema for a provider, creating default if not found"""
# Check if we have a specific schema method
method_name = f"get_{provider_name}_schema"
@@ -617,9 +614,7 @@ class StandardProviderSettings:
provider_name = attr_name[4:-7] # Remove 'get_' and '_schema'
builtin_providers.append(provider_name)
all_providers = list(
set(builtin_providers + list(cls._registered_schemas.keys()))
)
all_providers = list(set(builtin_providers + list(cls._registered_schemas.keys())))
return sorted(all_providers)
@classmethod
@@ -120,9 +120,7 @@ class StandardValidationRules:
except Exception:
return False
return ValidationRule(
name="valid_url", validator=validate_url, error_message=error_msg
)
return ValidationRule(name="valid_url", validator=validate_url, error_message=error_msg)
@staticmethod
def valid_ip_address(
@@ -155,9 +153,7 @@ class StandardValidationRules:
)
@staticmethod
def in_choices(
choices: List[Any], error_msg: Optional[str] = None
) -> ValidationRule:
def in_choices(choices: List[Any], error_msg: Optional[str] = None) -> ValidationRule:
"""Rule to ensure value is in predefined choices"""
if error_msg is None:
error_msg = f"Value must be one of: {', '.join(map(str, choices))}"
@@ -203,8 +199,7 @@ class SettingValue:
self.add_validation_rule(
ValidationRule(
name="is_integer",
validator=lambda x: isinstance(x, int)
or (isinstance(x, str) and x.isdigit()),
validator=lambda x: isinstance(x, int) or (isinstance(x, str) and x.isdigit()),
error_message="Value must be an integer",
)
)
@@ -279,16 +274,10 @@ class SettingValue:
return None
try:
if (
self.setting_type == SettingType.STRING
or self.setting_type == SettingType.PASSWORD
):
if self.setting_type == SettingType.STRING or self.setting_type == SettingType.PASSWORD:
return str(value)
elif (
self.setting_type == SettingType.INTEGER
or self.setting_type == SettingType.PORT
):
elif self.setting_type == SettingType.INTEGER or self.setting_type == SettingType.PORT:
if isinstance(value, int):
return value
elif isinstance(value, str) and value.isdigit():
@@ -355,9 +344,7 @@ class SettingValue:
def remove_validation_rule(self, rule_name: str) -> bool:
"""Remove a validation rule by name"""
original_length = len(self.validation_rules)
self.validation_rules = [
r for r in self.validation_rules if r.name != rule_name
]
self.validation_rules = [r for r in self.validation_rules if r.name != rule_name]
return len(self.validation_rules) < original_length
def validate(self) -> tuple[bool, List[str]]:
@@ -369,9 +356,7 @@ class SettingValue:
"""
if self.current_value is None:
# Check if this setting is required (has not_empty rule)
has_required_rule = any(
rule.name == "not_empty" for rule in self.validation_rules
)
has_required_rule = any(rule.name == "not_empty" for rule in self.validation_rules)
if has_required_rule:
return False, ["Value is required"]
else:
@@ -508,9 +493,7 @@ class SettingValueBuilder:
self, min_val: Union[int, float], max_val: Union[int, float]
) -> "SettingValueBuilder":
"""Add numeric range validation"""
self._validation_rules.append(
StandardValidationRules.numeric_range(min_val, max_val)
)
self._validation_rules.append(StandardValidationRules.numeric_range(min_val, max_val))
return self
def custom_validation(self, rule: ValidationRule) -> "SettingValueBuilder":
@@ -58,9 +58,7 @@ class ProviderEnableManager:
logger.debug(f"ProviderEnableManager: Kodi bridge not available: {e}")
self._kodi_bridge = None
except Exception as e:
logger.error(
f"ProviderEnableManager: Error initializing Kodi bridge: {e}"
)
logger.error(f"ProviderEnableManager: Error initializing Kodi bridge: {e}")
self._kodi_bridge = None
def _load_file(self, force_reload: bool = False) -> Dict[str, Any]:
@@ -76,11 +74,7 @@ class ProviderEnableManager:
current_time = time.time()
# Check cache first
if (
not force_reload
and self._cache
and (current_time - self._cache_time) < self.CACHE_TTL
):
if not force_reload and self._cache and (current_time - self._cache_time) < self.CACHE_TTL:
return self._cache
# Initialize default structure
@@ -100,9 +94,7 @@ class ProviderEnableManager:
data = self.vfs.read_json(self.DEFAULT_FILENAME)
if not isinstance(data, dict):
logger.warning(
f"ProviderEnableManager: Invalid file format, using defaults"
)
logger.warning(f"ProviderEnableManager: Invalid file format, using defaults")
self._cache = default_data
self._cache_time = current_time
return default_data
@@ -139,9 +131,7 @@ class ProviderEnableManager:
return data
except Exception as e:
logger.error(
f"ProviderEnableManager: Error loading {self.DEFAULT_FILENAME}: {e}"
)
logger.error(f"ProviderEnableManager: Error loading {self.DEFAULT_FILENAME}: {e}")
self._cache = default_data
self._cache_time = current_time
return default_data
@@ -173,16 +163,12 @@ class ProviderEnableManager:
f"ProviderEnableManager: Saved {len(data.get('providers', {}))} providers to file"
)
else:
logger.error(
f"ProviderEnableManager: Failed to write to {self.DEFAULT_FILENAME}"
)
logger.error(f"ProviderEnableManager: Failed to write to {self.DEFAULT_FILENAME}")
return success
except Exception as e:
logger.error(
f"ProviderEnableManager: Error saving {self.DEFAULT_FILENAME}: {e}"
)
logger.error(f"ProviderEnableManager: Error saving {self.DEFAULT_FILENAME}: {e}")
return False
def _get_kodi_enabled_status(self, provider_name: str) -> Optional[bool]:
@@ -228,9 +214,7 @@ class ProviderEnableManager:
return enabled
# No Kodi setting found
logger.debug(
f"ProviderEnableManager: No Kodi setting found for {provider_name}"
)
logger.debug(f"ProviderEnableManager: No Kodi setting found for {provider_name}")
return None
except Exception as e:
@@ -293,9 +277,7 @@ class ProviderEnableManager:
return EnableSource.DEFAULT
def set_provider_enabled(
self, provider_name: str, enabled: bool
) -> Tuple[bool, str]:
def set_provider_enabled(self, provider_name: str, enabled: bool) -> Tuple[bool, str]:
"""
Set enabled status for a provider (writes to file only).
@@ -400,9 +382,7 @@ class ProviderEnableManager:
source = self.get_enabled_source(provider_name)
if source == EnableSource.KODI:
message = (
f"Cannot delete setting for '{provider_name}' - controlled by Kodi"
)
message = f"Cannot delete setting for '{provider_name}' - controlled by Kodi"
logger.warning(f"ProviderEnableManager: {message}")
return False, message
@@ -515,17 +495,13 @@ class ProviderEnableManager:
# Kodi has explicit setting, migrate to file
data["providers"][provider_name] = kodi_enabled
migrated_count += 1
logger.debug(
f"ProviderEnableManager: Migrated {provider_name}={kodi_enabled}"
)
logger.debug(f"ProviderEnableManager: Migrated {provider_name}={kodi_enabled}")
if migrated_count > 0:
# Save migrated data
success = self._save_file(data)
if success:
message = (
f"Migrated {migrated_count} provider settings from Kodi to file"
)
message = f"Migrated {migrated_count} provider settings from Kodi to file"
logger.info(f"ProviderEnableManager: {message}")
return True, message, migrated_count
else:
@@ -12,8 +12,7 @@ from pathlib import Path
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
from ..auth.credential_manager import CredentialManager
from ..auth.credentials import (BaseCredentials, ClientCredentials,
UserPasswordCredentials)
from ..auth.credentials import BaseCredentials, ClientCredentials, UserPasswordCredentials
from ..auth.session_manager import SessionManager
from ..models.proxy_models import ProxyConfig
from ..network.proxy_manager import ProxyConfigManager
@@ -31,12 +30,8 @@ class ProviderRegistration:
registered_at: float = field(default_factory=time.time)
is_active: bool = True
settings_schema: Optional[Dict[str, Any]] = None
supports_countries: bool = (
False # NEW: Flag if provider supports country-specific settings
)
available_countries: List[str] = field(
default_factory=list
) # NEW: List of supported countries
supports_countries: bool = False # NEW: Flag if provider supports country-specific settings
available_countries: List[str] = field(default_factory=list) # NEW: List of supported countries
# Maintain backward compatibility by keeping the original class name
@@ -49,9 +44,7 @@ class SettingsManager:
Now supports country-specific settings for multi-region providers.
"""
def __init__(
self, config_dir: Optional[str] = None, enable_kodi_integration: bool = True
):
def __init__(self, config_dir: Optional[str] = None, enable_kodi_integration: bool = True):
"""
Initialize unified settings manager
@@ -112,9 +105,7 @@ class SettingsManager:
# Auto-register any unregistered providers
for provider_name, countries in detected_providers.items():
if not self.is_provider_registered(provider_name):
logger.info(
f"Auto-registering provider '{provider_name}' from Kodi"
)
logger.info(f"Auto-registering provider '{provider_name}' from Kodi")
self.register_provider(
provider_name,
supports_countries=bool(countries),
@@ -122,9 +113,7 @@ class SettingsManager:
)
else:
# Provider already registered - check if we need to upgrade to multi-country
if countries and not self.provider_supports_countries(
provider_name
):
if countries and not self.provider_supports_countries(provider_name):
logger.info(
f"Auto-upgrading existing provider '{provider_name}' to multi-country: {countries}"
)
@@ -144,9 +133,7 @@ class SettingsManager:
)
for country in available_countries:
logger.info(
f"Syncing credentials for {provider_name} ({country})..."
)
logger.info(f"Syncing credentials for {provider_name} ({country})...")
self._sync_credentials_from_kodi(provider_name, country)
logger.info(f"Syncing proxy for {provider_name} ({country})...")
@@ -159,9 +146,7 @@ class SettingsManager:
logger.info(f"Syncing proxy for {provider_name}...")
self._sync_proxy_from_kodi(provider_name)
logger.info(
f"Initialized SettingsManager with config dir: {self.config_dir_path}"
)
logger.info(f"Initialized SettingsManager with config dir: {self.config_dir_path}")
def _setup_kodi_integration(self) -> bool:
"""Setup Kodi integration, return True if successful"""
@@ -205,12 +190,8 @@ class SettingsManager:
registered_at=provider_info.get("registered_at", time.time()),
is_active=provider_info.get("is_active", True),
settings_schema=provider_info.get("settings_schema"),
supports_countries=provider_info.get(
"supports_countries", False
),
available_countries=provider_info.get(
"available_countries", []
),
supports_countries=provider_info.get("supports_countries", False),
available_countries=provider_info.get("available_countries", []),
)
self._registered_providers[provider_name] = registration
logger.debug(
@@ -220,9 +201,7 @@ class SettingsManager:
except Exception as e:
logger.error(f"Error loading provider {provider_name}: {e}")
logger.info(
f"Loaded {len(self._registered_providers)} provider registrations"
)
logger.info(f"Loaded {len(self._registered_providers)} provider registrations")
except Exception as e:
logger.error(f"Error loading unified settings configuration: {e}")
@@ -275,22 +254,16 @@ class SettingsManager:
BaseCredentials instance or None
"""
country_str = f" (country: {country})" if country else ""
logger.debug(
f"SettingsManager: Loading credentials for '{provider_name}{country_str}'"
)
logger.debug(f"SettingsManager: Loading credentials for '{provider_name}{country_str}'")
# Try file-based credentials first
credentials = self.credential_manager.load_credentials(provider_name, country)
logger.debug(f"SettingsManager: File credentials result: {type(credentials)}")
if credentials:
logger.debug(
f"SettingsManager: File credentials type: {credentials.credential_type}"
)
logger.debug(f"SettingsManager: File credentials type: {credentials.credential_type}")
if hasattr(credentials, "username"):
logger.debug(
f"SettingsManager: File credentials username: {credentials.username}"
)
logger.debug(f"SettingsManager: File credentials username: {credentials.username}")
else:
logger.debug(
f"SettingsManager: File credentials class: {credentials.__class__.__name__}"
@@ -301,21 +274,13 @@ class SettingsManager:
)
# If no file credentials and Kodi is available, try syncing from Kodi
if (
not credentials
and self.kodi_bridge
and self.kodi_bridge.is_kodi_environment()
):
if not credentials and self.kodi_bridge and self.kodi_bridge.is_kodi_environment():
logger.debug(
f"SettingsManager: No file credentials for {provider_name}{country_str}, trying Kodi sync"
)
if self._sync_credentials_from_kodi(provider_name, country):
credentials = self.credential_manager.load_credentials(
provider_name, country
)
logger.debug(
f"SettingsManager: After Kodi sync, credentials: {type(credentials)}"
)
credentials = self.credential_manager.load_credentials(provider_name, country)
logger.debug(f"SettingsManager: After Kodi sync, credentials: {type(credentials)}")
if credentials:
logger.debug(
f"SettingsManager: Synced credentials type: {credentials.credential_type}"
@@ -344,15 +309,11 @@ class SettingsManager:
True if successful, False otherwise
"""
# Save to file first
file_success = self.credential_manager.save_credentials(
provider_name, credentials, country
)
file_success = self.credential_manager.save_credentials(provider_name, credentials, country)
if not file_success:
country_str = f" (country: {country})" if country else ""
logger.error(
f"Failed to save credentials to file for {provider_name}{country_str}"
)
logger.error(f"Failed to save credentials to file for {provider_name}{country_str}")
return False
# Sync to Kodi if available and in Kodi environment
@@ -361,9 +322,7 @@ class SettingsManager:
provider_name, credentials, country
)
if not kodi_success:
logger.warning(
f"Failed to sync credentials to Kodi for {provider_name}"
)
logger.warning(f"Failed to sync credentials to Kodi for {provider_name}")
# Don't fail the entire operation if Kodi sync fails
return True
@@ -468,22 +427,16 @@ class SettingsManager:
success = self._save_configuration()
if success:
country_info = (
f" (supports countries: {available_countries})"
if supports_countries
else ""
)
logger.info(
f"Successfully registered provider: {provider_name}{country_info}"
f" (supports countries: {available_countries})" if supports_countries else ""
)
logger.info(f"Successfully registered provider: {provider_name}{country_info}")
return success
except Exception as e:
logger.error(f"Error registering provider {provider_name}: {e}")
return False
def unregister_provider(
self, provider_name: str, cleanup_data: bool = False
) -> bool:
def unregister_provider(self, provider_name: str, cleanup_data: bool = False) -> bool:
"""
Unregister a provider and optionally clean up its data
@@ -505,9 +458,7 @@ class SettingsManager:
reg = self._registered_providers[provider_name]
if reg.supports_countries:
for country in reg.available_countries:
self.credential_manager.delete_credentials(
provider_name, country
)
self.credential_manager.delete_credentials(provider_name, country)
self.session_manager.clear_session(provider_name, country)
self.proxy_manager.remove_proxy_config(provider_name, country)
else:
@@ -531,9 +482,7 @@ class SettingsManager:
def list_registered_providers(self) -> List[str]:
"""Get list of all registered providers"""
return [
name for name, reg in self._registered_providers.items() if reg.is_active
]
return [name for name, reg in self._registered_providers.items() if reg.is_active]
def is_provider_registered(self, provider_name: str) -> bool:
"""Check if a provider is registered"""
@@ -632,9 +581,7 @@ class SettingsManager:
"kodi_integration": {
"enabled": self.enable_kodi_integration,
"environment": (
self.kodi_bridge.is_kodi_environment()
if self.kodi_bridge
else False
self.kodi_bridge.is_kodi_environment() if self.kodi_bridge else False
),
},
}
@@ -659,9 +606,7 @@ class SettingsManager:
return status
def is_provider_ready(
self, provider_name: str, country: Optional[str] = None
) -> bool:
def is_provider_ready(self, provider_name: str, country: Optional[str] = None) -> bool:
"""
Check if provider is fully configured and ready to use
@@ -679,9 +624,9 @@ class SettingsManager:
# Provider is ready if it has valid credentials
credentials_status = status.get("credentials", {})
return credentials_status.get(
"has_credentials", False
) and credentials_status.get("credentials_valid", False)
return credentials_status.get("has_credentials", False) and credentials_status.get(
"credentials_valid", False
)
def get_ready_providers(self, country: Optional[str] = None) -> List[str]:
"""
@@ -715,9 +660,7 @@ class SettingsManager:
True if successful or no sync needed, False on error
"""
country_str = f" (country: {country})" if country else ""
logger.debug(
f"SettingsManager: Attempting Kodi sync for {provider_name}{country_str}"
)
logger.debug(f"SettingsManager: Attempting Kodi sync for {provider_name}{country_str}")
if not self.kodi_bridge or not self.kodi_bridge.is_kodi_environment():
logger.debug(
@@ -726,12 +669,8 @@ class SettingsManager:
return False
try:
kodi_credentials = self.kodi_bridge.read_credentials_from_kodi(
provider_name, country
)
logger.debug(
f"SettingsManager: Kodi credentials result: {type(kodi_credentials)}"
)
kodi_credentials = self.kodi_bridge.read_credentials_from_kodi(provider_name, country)
logger.debug(f"SettingsManager: Kodi credentials result: {type(kodi_credentials)}")
if not kodi_credentials:
logger.debug(
@@ -746,14 +685,10 @@ class SettingsManager:
return True # Not an error
# Check if different from current file credentials
file_credentials = self.credential_manager.load_credentials(
provider_name, country
)
file_credentials = self.credential_manager.load_credentials(provider_name, country)
# ← FIX: Only skip if BOTH exist AND are equal
if file_credentials and self._credentials_equal(
kodi_credentials, file_credentials
):
if file_credentials and self._credentials_equal(kodi_credentials, file_credentials):
logger.debug(
f"SettingsManager: Credentials already in sync for {provider_name}{country_str}"
)
@@ -769,9 +704,7 @@ class SettingsManager:
logger.debug(f"SettingsManager: Kodi sync save result: {success}")
if success:
logger.info(
f"Synced credentials from Kodi for {provider_name}{country_str}"
)
logger.info(f"Synced credentials from Kodi for {provider_name}{country_str}")
return success
except Exception as e:
@@ -809,13 +742,9 @@ class SettingsManager:
if countries:
# Multi-country provider
for country in countries:
cred_success = self._sync_credentials_from_kodi(
provider_name, country
)
cred_success = self._sync_credentials_from_kodi(provider_name, country)
proxy_success = self._sync_proxy_from_kodi(provider_name, country)
results[f"{provider_name}_{country}"] = (
cred_success or proxy_success
)
results[f"{provider_name}_{country}"] = cred_success or proxy_success
else:
# Single provider
cred_success = self._sync_credentials_from_kodi(provider_name)
@@ -853,32 +782,20 @@ class SettingsManager:
proxy_config = self.proxy_manager.get_proxy_config(provider_name, country)
# If no proxy config and Kodi available, sync from Kodi
if (
not proxy_config
and self.kodi_bridge
and self.kodi_bridge.is_kodi_environment()
):
if not proxy_config and self.kodi_bridge and self.kodi_bridge.is_kodi_environment():
if self._sync_proxy_from_kodi(provider_name, country):
proxy_config = self.proxy_manager.get_proxy_config(
provider_name, country
)
proxy_config = self.proxy_manager.get_proxy_config(provider_name, country)
return proxy_config
def _sync_proxy_from_kodi(
self, provider_name: str, country: Optional[str] = None
) -> bool:
def _sync_proxy_from_kodi(self, provider_name: str, country: Optional[str] = None) -> bool:
"""Sync proxy config from Kodi to proxy manager"""
if not self.kodi_bridge:
return False
kodi_proxy_config = self.kodi_bridge.read_proxy_config_from_kodi(
provider_name, country
)
kodi_proxy_config = self.kodi_bridge.read_proxy_config_from_kodi(provider_name, country)
if kodi_proxy_config:
return self.proxy_manager.set_proxy_config(
provider_name, kodi_proxy_config, country
)
return self.proxy_manager.set_proxy_config(provider_name, kodi_proxy_config, country)
return False
def set_provider_proxy(
@@ -888,29 +805,21 @@ class SettingsManager:
country: Optional[str] = None,
) -> bool:
"""Set proxy configuration for a provider with Kodi sync"""
success = self.proxy_manager.set_proxy_config(
provider_name, proxy_config, country
)
success = self.proxy_manager.set_proxy_config(provider_name, proxy_config, country)
# Also sync to Kodi if available
if success and self.kodi_bridge and self.kodi_bridge.is_kodi_environment():
self.kodi_bridge.write_proxy_config_to_kodi(
provider_name, proxy_config, country
)
self.kodi_bridge.write_proxy_config_to_kodi(provider_name, proxy_config, country)
return success
def remove_provider_proxy(
self, provider_name: str, country: Optional[str] = None
) -> bool:
def remove_provider_proxy(self, provider_name: str, country: Optional[str] = None) -> bool:
"""Remove proxy configuration for a provider"""
return self.proxy_manager.remove_proxy_config(provider_name, country)
# ============= Session Management =============
def clear_provider_session(
self, provider_name: str, country: Optional[str] = None
) -> bool:
def clear_provider_session(self, provider_name: str, country: Optional[str] = None) -> bool:
"""Clear all session data for a provider"""
return self.session_manager.clear_session(provider_name, country)
@@ -990,8 +899,7 @@ class SettingsManager:
if reset_session:
self.session_manager.clear_session(provider_name, country)
logger.info(
f"Reset session for {provider_name}"
+ (f" ({country})" if country else "")
f"Reset session for {provider_name}" + (f" ({country})" if country else "")
)
if reset_proxy:
@@ -1054,9 +962,7 @@ class SettingsManager:
if country is None and self.provider_supports_countries(provider_name):
countries_data = {}
for ctry in self.get_provider_countries(provider_name):
countries_data[ctry] = self.export_provider_settings(
provider_name, ctry
)
countries_data[ctry] = self.export_provider_settings(provider_name, ctry)
export_data["countries"] = countries_data
return export_data
@@ -1075,9 +981,7 @@ class SettingsManager:
timestamp = int(time.time())
# Use config_path for backward compatibility
if self.config_path:
export_path = str(
self.config_path / f"settings_export_{timestamp}.json"
)
export_path = str(self.config_path / f"settings_export_{timestamp}.json")
else:
export_path = f"settings_export_{timestamp}.json"
@@ -1088,9 +992,7 @@ class SettingsManager:
"source": "unified_settings_manager_v2.1_country_aware",
},
"system_info": {
"config_dir": (
str(self.config_path) if self.config_path else "default_vfs_path"
),
"config_dir": (str(self.config_path) if self.config_path else "default_vfs_path"),
"kodi_integration_enabled": self.enable_kodi_integration,
"registered_provider_count": len(self._registered_providers),
},
@@ -1099,9 +1001,7 @@ class SettingsManager:
# Export each provider
for provider_name in self.list_registered_providers():
export_data["providers"][provider_name] = self.export_provider_settings(
provider_name
)
export_data["providers"][provider_name] = self.export_provider_settings(provider_name)
# Write to file - use standard filesystem for exports (not VFS)
with open(export_path, "w", encoding="utf-8") as f:
@@ -1140,18 +1040,14 @@ class SettingsManager:
"active": len(self.list_registered_providers()),
"ready": len(self.get_ready_providers()),
"country_aware": sum(
1
for reg in self._registered_providers.values()
if reg.supports_countries
1 for reg in self._registered_providers.values() if reg.supports_countries
),
},
"kodi_integration": {
"enabled": self.enable_kodi_integration,
"bridge_available": self.kodi_bridge is not None,
"in_kodi_environment": (
self.kodi_bridge.is_kodi_environment()
if self.kodi_bridge
else False
self.kodi_bridge.is_kodi_environment() if self.kodi_bridge else False
),
},
"storage": {
@@ -1160,9 +1056,7 @@ class SettingsManager:
},
}
def debug_provider(
self, provider_name: str, country: Optional[str] = None
) -> Dict[str, Any]:
def debug_provider(self, provider_name: str, country: Optional[str] = None) -> Dict[str, Any]:
"""
Get detailed debug information for a provider
@@ -1212,8 +1106,7 @@ class SettingsManager:
# Component manager status
debug_info["components"] = {
"credential_manager": {
"has_credentials": provider_name
in self.credential_manager.list_providers()
"has_credentials": provider_name in self.credential_manager.list_providers()
},
"session_manager": {
"has_session": self.session_manager.load_session(provider_name, country)
@@ -1284,9 +1177,7 @@ class SettingsManager:
# Get countries with session data
return self.session_manager.get_all_countries(provider_name)
def migrate_to_country_structure(
self, provider_name: str, default_country: str
) -> bool:
def migrate_to_country_structure(self, provider_name: str, default_country: str) -> bool:
"""
Migrate a non-country provider to country-aware structure
@@ -1323,15 +1214,11 @@ class SettingsManager:
logger.debug(f"Migrated credentials to {default_country}")
if old_session:
self.session_manager.save_session(
provider_name, old_session, default_country
)
self.session_manager.save_session(provider_name, old_session, default_country)
logger.debug(f"Migrated session to {default_country}")
if old_proxy:
self.proxy_manager.set_proxy_config(
provider_name, old_proxy, default_country
)
self.proxy_manager.set_proxy_config(provider_name, old_proxy, default_country)
logger.debug(f"Migrated proxy to {default_country}")
# Clear old non-country data
@@ -1395,9 +1282,7 @@ class SettingsManager:
'token_type': 'Bearer'
}, 'de')
"""
return self.session_manager.save_scoped_token(
provider_name, scope, token_data, country
)
return self.session_manager.save_scoped_token(provider_name, scope, token_data, country)
def load_scoped_token(
self, provider_name: str, scope: str, country: Optional[str] = None
@@ -1443,9 +1328,7 @@ class SettingsManager:
"""
return self.session_manager.clear_scoped_token(provider_name, scope, country)
def list_scoped_tokens(
self, provider_name: str, country: Optional[str] = None
) -> List[str]:
def list_scoped_tokens(self, provider_name: str, country: Optional[str] = None) -> List[str]:
"""
List all available token scopes for a provider
@@ -1529,13 +1412,9 @@ class SettingsManager:
status["access_token_valid"] = time.time() < expires_at
# yo_digital style (separate access_token_expires_in)
elif (
"access_token_expires_in" in token_data
and "access_token_issued_at" in token_data
):
elif "access_token_expires_in" in token_data and "access_token_issued_at" in token_data:
expires_at = (
token_data["access_token_issued_at"]
+ token_data["access_token_expires_in"]
token_data["access_token_issued_at"] + token_data["access_token_expires_in"]
)
status["access_token_expires_at"] = expires_at
status["access_token_valid"] = time.time() < expires_at
@@ -1546,13 +1425,9 @@ class SettingsManager:
if "refresh_token" in token_data:
status["has_refresh_token"] = True
if (
"refresh_token_expires_in" in token_data
and "refresh_token_issued_at" in token_data
):
if "refresh_token_expires_in" in token_data and "refresh_token_issued_at" in token_data:
expires_at = (
token_data["refresh_token_issued_at"]
+ token_data["refresh_token_expires_in"]
token_data["refresh_token_issued_at"] + token_data["refresh_token_expires_in"]
)
status["refresh_token_expires_at"] = expires_at
status["refresh_token_valid"] = time.time() < expires_at
@@ -1561,9 +1436,7 @@ class SettingsManager:
return status
def is_provider_enabled(
self, provider_name: str, country: Optional[str] = None
) -> bool:
def is_provider_enabled(self, provider_name: str, country: Optional[str] = None) -> bool:
"""Check if provider is enabled (using ProviderEnableManager)"""
# Import here to avoid circular imports
from .provider_enable_manager import ProviderEnableManager
@@ -1647,15 +1520,11 @@ class SettingsManager:
return False, f"Credential validation failed: missing required fields"
# Save directly to file (bypass Kodi sync)
success = self.credential_manager.save_credentials(
provider, credentials, country
)
success = self.credential_manager.save_credentials(provider, credentials, country)
if success:
country_str = f" ({country})" if country else ""
logger.info(
f"Successfully saved credentials from API for {provider}{country_str}"
)
logger.info(f"Successfully saved credentials from API for {provider}{country_str}")
return True, "Credentials saved successfully"
else:
return False, "Failed to save credentials to file"
@@ -1689,9 +1558,7 @@ class SettingsManager:
)
# Check for client credentials
has_client_id = (
"client_id" in credentials_data and credentials_data["client_id"]
)
has_client_id = "client_id" in credentials_data and credentials_data["client_id"]
has_client_secret = (
"client_secret" in credentials_data and credentials_data["client_secret"]
)
@@ -1744,15 +1611,11 @@ class SettingsManager:
return False, "Proxy validation failed: invalid host, port, or timeout"
# Save directly to file (bypass Kodi sync)
success = self.proxy_manager.set_proxy_config(
provider, proxy_config, country
)
success = self.proxy_manager.set_proxy_config(provider, proxy_config, country)
if success:
country_str = f" ({country})" if country else ""
logger.info(
f"Successfully saved proxy config from API for {provider}{country_str}"
)
logger.info(f"Successfully saved proxy config from API for {provider}{country_str}")
return True, "Proxy configuration saved successfully"
else:
return False, "Failed to save proxy configuration to file"
@@ -1772,8 +1635,7 @@ class SettingsManager:
Returns:
ProxyConfig instance or None if invalid
"""
from ..models.proxy_models import (ProxyAuth, ProxyConfig, ProxyScope,
ProxyType)
from ..models.proxy_models import ProxyAuth, ProxyConfig, ProxyScope, ProxyType
# Required fields
if "host" not in proxy_data or not proxy_data["host"]:
@@ -1829,9 +1691,7 @@ class SettingsManager:
verify_ssl=verify_ssl,
)
def delete_provider_credentials_from_api(
self, provider_name: str
) -> Tuple[bool, str]:
def delete_provider_credentials_from_api(self, provider_name: str) -> Tuple[bool, str]:
"""
Delete credentials via API request
@@ -1862,9 +1722,7 @@ class SettingsManager:
return False, "Failed to delete credentials"
except Exception as e:
logger.error(
f"Error deleting credentials from API for {provider_name}: {e}"
)
logger.error(f"Error deleting credentials from API for {provider_name}: {e}")
return False, f"Internal error: {str(e)}"
def delete_provider_proxy_from_api(self, provider_name: str) -> Tuple[bool, str]:
@@ -1898,9 +1756,7 @@ class SettingsManager:
return False, "Failed to delete proxy configuration"
except Exception as e:
logger.error(
f"Error deleting proxy config from API for {provider_name}: {e}"
)
logger.error(f"Error deleting proxy config from API for {provider_name}: {e}")
return False, f"Internal error: {str(e)}"
@@ -17,9 +17,7 @@ class SubscriptionOperations:
self.registry = registry
logger.debug("SubscriptionOperations: Initialized")
def get_subscription_status(
self, provider_name: str, **kwargs
) -> Optional[UserSubscription]:
def get_subscription_status(self, provider_name: str, **kwargs) -> Optional[UserSubscription]:
"""Get subscription status for a provider."""
provider = self.registry.get_provider(provider_name)
if not provider:
@@ -37,9 +35,7 @@ class SubscriptionOperations:
logger.warning(f"Error getting subscription for '{provider_name}': {e}")
return None
def get_subscribed_channels(
self, provider_name: str, **kwargs
) -> List[StreamingChannel]:
def get_subscribed_channels(self, provider_name: str, **kwargs) -> List[StreamingChannel]:
"""Get subscribed channels."""
provider = self.registry.get_provider(provider_name)
if not provider:
@@ -47,17 +43,13 @@ class SubscriptionOperations:
try:
channels = provider.get_subscribed_channels(**kwargs)
logger.info(
f"Got {len(channels)} subscribed channels from '{provider_name}'"
)
logger.info(f"Got {len(channels)} subscribed channels from '{provider_name}'")
return channels
except Exception as e:
logger.error(f"Error getting subscribed channels: {e}")
return provider.get_channels(**kwargs)
def get_available_packages(
self, provider_name: str, **kwargs
) -> List[SubscriptionPackage]:
def get_available_packages(self, provider_name: str, **kwargs) -> List[SubscriptionPackage]:
"""Get available subscription packages."""
provider = self.registry.get_provider(provider_name)
if not provider:
@@ -71,9 +63,7 @@ class SubscriptionOperations:
logger.warning(f"Error getting packages for '{provider_name}': {e}")
return []
def is_channel_accessible(
self, provider_name: str, channel_id: str, **kwargs
) -> bool:
def is_channel_accessible(self, provider_name: str, channel_id: str, **kwargs) -> bool:
"""Check if channel is accessible with current subscription."""
provider = self.registry.get_provider(provider_name)
if not provider:
@@ -88,9 +88,7 @@ class ConsoleNotificationAdapter(NotificationInterface):
print("=" * 70)
print()
logger.info(
f"Remote login started: code={login_code}, expires_in={expires_in}s"
)
logger.info(f"Remote login started: code={login_code}, expires_in={expires_in}s")
logger.info(f"QR target URL: {qr_target_url}")
return NotificationResult.CONTINUE
@@ -286,9 +286,7 @@ class QRCodeDialog:
if self.polling_thread:
remaining = self.polling_thread.get_remaining_time()
if self.time_label:
self.time_label.setLabel(
f"Time remaining: {self._format_time(remaining)}"
)
self.time_label.setLabel(f"Time remaining: {self._format_time(remaining)}")
# Check if auth completed
if self.polling_thread.auth_completed:
@@ -303,13 +301,9 @@ class QRCodeDialog:
# Check for errors
if self.polling_thread.error:
logger.error(
f"Monitor: Polling error: {self.polling_thread.error}"
)
logger.error(f"Monitor: Polling error: {self.polling_thread.error}")
if self.status_label:
self.status_label.setLabel(
"[COLOR red]Authentication failed[/COLOR]"
)
self.status_label.setLabel("[COLOR red]Authentication failed[/COLOR]")
time.sleep(2)
self.close_dialog()
break
@@ -41,14 +41,8 @@ class NotificationFactory:
NotificationInterface: Appropriate adapter for current environment
"""
# Return cached adapter if available and no http_manager change
if (
cls._cached_adapter is not None
and force_environment is None
and http_manager is None
):
logger.debug(
f"Using cached notification adapter: {cls._environment_detected}"
)
if cls._cached_adapter is not None and force_environment is None and http_manager is None:
logger.debug(f"Using cached notification adapter: {cls._environment_detected}")
return cls._cached_adapter
# Detect environment
@@ -89,9 +83,7 @@ class NotificationFactory:
except ImportError:
# If import fails, we're standalone
logger.debug(
"Kodi modules not available - using console notification adapter"
)
logger.debug("Kodi modules not available - using console notification adapter")
return "console"
@classmethod
@@ -41,9 +41,7 @@ def generate_qr_code_png(data: str, size: int = 512) -> Optional[bytes]:
qr.make(fit=True)
# Generate image using pure Python PNG backend
img = qr.make_image(
image_factory=PyPNGImage, fill_color="black", back_color="white"
)
img = qr.make_image(image_factory=PyPNGImage, fill_color="black", back_color="white")
# Convert to PNG bytes
buffer = io.BytesIO()
@@ -83,22 +83,16 @@ class EnvironmentManager:
# Get settings
default_country = self._addon.getSetting("default_country")
self._config["default_country"] = (
str(default_country) if default_country else "DE"
)
self._config["default_country"] = str(default_country) if default_country else "DE"
server_port = self._addon.getSetting("server_port")
try:
self._config["server_port"] = (
int(str(server_port)) if server_port else 7777
)
self._config["server_port"] = int(str(server_port)) if server_port else 7777
except ValueError:
self._config["server_port"] = 7777
except Exception as init_error: # noqa: B902
print(
f"DEBUG: Exception type: {type(init_error).__name__}", file=sys.stderr
)
print(f"DEBUG: Exception type: {type(init_error).__name__}", file=sys.stderr)
print(f"DEBUG: Exception message: {str(init_error)}", file=sys.stderr)
# Log the error and fallback to standalone
self._log_init_error("Kodi initialization failed", init_error)
@@ -120,9 +114,7 @@ class EnvironmentManager:
self._config["addon_path"] = os.path.dirname(os.path.abspath(__file__))
# Default configuration paths
config_home = os.environ.get("XDG_CONFIG_HOME") or os.path.join(
str(Path.home()), ".config"
)
config_home = os.environ.get("XDG_CONFIG_HOME") or os.path.join(str(Path.home()), ".config")
self._config["config_dir"] = os.path.join(config_home, "ultimate-backend")
self._config["profile_path"] = self._config["config_dir"]
@@ -248,9 +240,7 @@ class EnvironmentManager:
else:
raise ImportError("get_configured_manager is not callable")
else:
raise ImportError(
"get_configured_manager not found in streaming_providers module"
)
raise ImportError("get_configured_manager not found in streaming_providers module")
except ImportError as manager_error:
# Log error using the logger once we have it
logger_instance = self.get_logger()
+1 -3
View File
@@ -50,9 +50,7 @@ class BaseLogger:
log_message += f" - {details}"
self.info(log_message)
def log_credential_event(
self, provider: str, event: str, details: str = ""
) -> None:
def log_credential_event(self, provider: str, event: str, details: str = "") -> None:
"""Log credential event"""
log_message = f"CRED [{provider}] {event}"
if details:
@@ -1,96 +1,166 @@
import base64
import re
from typing import List, Optional
from urllib.parse import quote, urljoin, urlparse
from ..models.drm_models import PSSHData
from ..models.drm_models import DRMSystem, PSSHData
from .logger import logger
class ManifestParser:
@staticmethod
def extract_pssh_from_manifest(
manifest_content: str, manifest_url: str = "", return_collection: bool = True
):
logger.debug("Starting PSSH extraction from manifest")
manifest_content: str,
manifest_url: str = "",
fallback_to_segments: bool = True,
segment_urls: List[str] = None,
) -> List[PSSHData]:
"""
DEPRECATED: Use extract_single_init_segment_url instead.
Kept for backwards compatibility.
"""
pssh_list = ManifestParser._extract_from_manifest_content(manifest_content)
# Check if this looks like a DASH manifest
is_dash = ("<MPD" in manifest_content) or ("mpd" in manifest_content.lower())
logger.debug(f"Manifest appears to be DASH format: {is_dash}")
if fallback_to_segments and segment_urls:
incomplete_pssh = [p for p in pssh_list if not p.pssh_box or not p.key_ids]
if incomplete_pssh:
segment_pssh = ManifestParser._extract_from_single_segment(
segment_urls[0], [p.system_id for p in incomplete_pssh]
)
return ManifestParser._merge_pssh_data(pssh_list, segment_pssh)
if not is_dash:
logger.debug("Not a DASH manifest, skipping PSSH extraction")
return []
try:
pssh_list = ManifestParser._extract_with_regex(manifest_content)
logger.debug(f"Found {len(pssh_list)} potential PSSH entries")
valid_pssh = []
for pssh in pssh_list:
if pssh.pssh_box and pssh.system_id:
valid_pssh.append(pssh)
logger.debug(f"Valid PSSH found - System ID: {pssh.system_id}")
else:
logger.debug("Invalid PSSH entry skipped")
return valid_pssh
except Exception as e:
logger.error(f"Error in PSSH extraction: {str(e)}")
return []
return pssh_list
@staticmethod
def _extract_with_regex(mpd_content: str):
"""Simplified PSSH extraction using regular expressions with debug logging"""
logger.debug("Starting regex PSSH extraction")
def _extract_from_manifest_content(manifest_content: str) -> List[PSSHData]:
"""Extract PSSH and DRM systems from manifest content"""
# Try regex extraction first (handles PSSH boxes with KIDs)
pssh_list = ManifestParser._extract_with_regex(manifest_content)
if pssh_list:
return pssh_list
pssh_dict = {} # Use dict to automatically handle deduplication
global_key_ids = [] # Collect all KIDs from the manifest
# Fallback: extract DRM systems from schemeIdUri only
drm_systems_found = set()
result = []
# Regex patterns
pssh_pattern = r"<(?:cenc:)?pssh[^>]*>([^<]+)</(?:cenc:)?pssh>"
default_kid_pattern = r'(?:cenc:)?default_KID="([^"]+)"'
system_id_pattern = r'schemeIdUri="urn:uuid:([^"]+)"'
# More efficient: compile regex once
cp_pattern = re.compile(
r'<ContentProtection[^>]*schemeIdUri="urn:uuid:([^"]+)"[^>]*>',
re.IGNORECASE,
)
logger.debug("Searching for ContentProtection blocks")
for match in cp_pattern.finditer(manifest_content):
system_id = match.group(1).lower()
# Skip mp4protection scheme
if "mp4protection" in manifest_content[max(0, match.start() - 100) : match.start()]:
continue
drm_system = DRMSystem.from_uuid(system_id)
if drm_system and system_id not in drm_systems_found:
drm_systems_found.add(system_id)
result.append(
PSSHData(
system_id=system_id,
pssh_box="", # Empty - PSSH in segments
key_ids=[],
source="manifest_scheme_only",
)
)
logger.debug(f"Found DRM system from schemeIdUri: {drm_system.value}")
return result
@staticmethod
def _extract_from_single_segment(
segment_url: str, expected_system_ids: List[str] = None
) -> List[PSSHData]:
"""Extract PSSH from a single MP4 segment"""
from .mp4_parser import MP4PSSHExtractor
try:
pssh_from_segment = MP4PSSHExtractor.extract_from_url(segment_url)
# Filter for expected DRM systems if provided
if expected_system_ids:
filtered_pssh = [p for p in pssh_from_segment if p.system_id in expected_system_ids]
if filtered_pssh:
logger.debug(f"Found {len(filtered_pssh)} PSSH boxes in segment")
return filtered_pssh
elif pssh_from_segment:
logger.debug(f"Found {len(pssh_from_segment)} PSSH boxes in segment")
return pssh_from_segment
except Exception as e:
logger.warning(f"Failed to extract PSSH from segment: {e}")
return []
@staticmethod
def _merge_pssh_data(
manifest_pssh: List[PSSHData], segment_pssh: List[PSSHData]
) -> List[PSSHData]:
"""Merge manifest and segment PSSH data"""
if not manifest_pssh:
return segment_pssh
if not segment_pssh:
return manifest_pssh
merged = []
segment_by_system = {p.system_id: p for p in segment_pssh}
for manifest_p in manifest_pssh:
if manifest_p.system_id in segment_by_system:
# Use segment data (complete)
merged.append(segment_by_system[manifest_p.system_id])
else:
# Keep manifest data (incomplete)
merged.append(manifest_p)
return merged
@staticmethod
def _extract_with_regex(mpd_content: str) -> List[PSSHData]:
"""Extract PSSH boxes using regex"""
pssh_dict = {}
global_key_ids = []
# Compile patterns once
pssh_pattern = re.compile(r"<(?:cenc:)?pssh[^>]*>([^<]+)</(?:cenc:)?pssh>")
default_kid_pattern = re.compile(r'(?:cenc:)?default_KID="([^"]+)"', re.IGNORECASE)
system_id_pattern = re.compile(r'schemeIdUri="urn:uuid:([^"]+)"', re.IGNORECASE)
# Find ContentProtection blocks efficiently
cp_blocks = re.findall(
r"<ContentProtection[^>]*>.*?</ContentProtection>", mpd_content, re.DOTALL
)
logger.debug(f"Found {len(cp_blocks)} ContentProtection blocks")
# First pass: collect all default KIDs from the entire manifest
# First pass: collect all default KIDs
for block in cp_blocks:
kid_match = re.search(default_kid_pattern, block)
kid_match = default_kid_pattern.search(block)
if kid_match:
clean_kid = kid_match.group(1).replace("-", "").lower()
if clean_kid not in global_key_ids:
global_key_ids.append(clean_kid)
logger.debug(f"Found global default KID: {clean_kid}")
# Second pass: extract PSSH data
for i, block in enumerate(cp_blocks, 1):
for block in cp_blocks:
try:
logger.debug(f"Processing block {i}/{len(cp_blocks)}")
system_id = None
# Extract system ID from schemeIdUri
scheme_match = re.search(system_id_pattern, block)
scheme_match = system_id_pattern.search(block)
if scheme_match:
system_id = scheme_match.group(1).lower()
logger.debug(f"Found system ID in schemeIdUri: {system_id}")
# Extract PSSH data
pssh_matches = re.findall(pssh_pattern, block)
logger.debug(f"Found {len(pssh_matches)} PSSH elements in block")
for pssh_b64 in pssh_matches:
for pssh_match in pssh_pattern.finditer(block):
pssh_b64 = pssh_match.group(1)
try:
logger.debug(
f"Processing PSSH (first 30 chars): {pssh_b64[:30]}..."
)
pssh_data = base64.b64decode(pssh_b64)
if len(pssh_data) >= 28:
# Extract system ID from PSSH if not found in schemeIdUri
# Extract system ID from PSSH if not found
if not system_id:
system_id_bytes = pssh_data[12:28]
system_id = "-".join(
@@ -102,29 +172,133 @@ class ManifestParser:
system_id_bytes[10:16].hex(),
]
)
logger.debug(
f"Extracted system ID from PSSH: {system_id}"
)
# Use PSSH box as key for deduplication
# Deduplicate by PSSH box content
if pssh_b64 not in pssh_dict:
pssh_dict[pssh_b64] = PSSHData(
system_id=system_id,
pssh_box=pssh_b64,
key_ids=global_key_ids.copy(), # Add all global KIDs to each PSSH
key_ids=global_key_ids.copy(),
source="manifest_pssh",
)
logger.debug("Successfully added new PSSH entry")
else:
logger.debug("PSSH already exists, skipping duplicate")
except Exception as e:
logger.error(f"Error decoding PSSH: {str(e)}")
logger.debug(f"Error decoding PSSH: {e}")
except Exception as e:
logger.error(f"Error processing ContentProtection block: {str(e)}")
logger.debug(f"Error processing ContentProtection block: {e}")
pssh_list = list(pssh_dict.values())
logger.debug(
f"Completed PSSH extraction, found {len(pssh_list)} unique entries"
return list(pssh_dict.values())
@staticmethod
def extract_single_init_segment_url(manifest_content: str, manifest_url: str) -> Optional[str]:
"""
Extract ONE init segment URL from DASH manifest.
Prioritizes video representations as they typically have the same DRM as audio.
"""
# Parse manifest base URL
parsed = urlparse(manifest_url)
manifest_base = f"{parsed.scheme}://{parsed.netloc}{'/'.join(parsed.path.split('/')[:-1])}"
if not manifest_base.endswith("/"):
manifest_base += "/"
# Extract BaseURL elements (can appear at multiple levels)
base_urls = re.findall(r"<BaseURL[^>]*>([^<]+)</BaseURL>", manifest_content)
# Build effective base URL
effective_base = manifest_base
for base_url in base_urls:
if base_url.startswith("http"):
effective_base = base_url
else:
effective_base = urljoin(effective_base, base_url)
if not effective_base.endswith("/"):
effective_base += "/"
logger.debug(f"Effective base URL: {effective_base}")
# Find SegmentTemplate with initialization attribute
# Prioritize video AdaptationSets
adaptation_sets = re.findall(
r"<AdaptationSet[^>]*>.*?</AdaptationSet>", manifest_content, re.DOTALL
)
return pssh_list
video_sets = []
audio_sets = []
for ad_set in adaptation_sets:
if 'contentType="video"' in ad_set or 'mimeType="video/' in ad_set:
video_sets.append(ad_set)
elif 'contentType="audio"' in ad_set or 'mimeType="audio/' in ad_set:
audio_sets.append(ad_set)
# Try video first, then audio
target_sets = video_sets + audio_sets
for ad_set in target_sets:
# Find SegmentTemplate initialization
seg_template_match = re.search(
r'<SegmentTemplate[^>]*initialization="([^"]+)"', ad_set, re.IGNORECASE
)
if not seg_template_match:
continue
init_template = seg_template_match.group(1)
logger.debug(f"Found init template: {init_template}")
# Find first Representation in this AdaptationSet
rep_match = re.search(r'<Representation[^>]*id="([^"]+)"', ad_set)
if not rep_match:
continue
rep_id = rep_match.group(1)
logger.debug(f"Using Representation ID: {rep_id}")
# Substitute template variables
init_url = init_template.replace("$RepresentationID$", rep_id)
# Handle other common template variables
init_url = init_url.replace("$Bandwidth$", "0")
init_url = init_url.replace("$Time$", "0")
init_url = init_url.replace("$Number$", "1")
# Construct full URL
if init_url.startswith("http"):
full_url = init_url
else:
# URL encode special characters in representation ID
# Split path and encode only the filename part
path_parts = init_url.split("/")
path_parts[-1] = quote(path_parts[-1], safe=".-_")
init_url = "/".join(path_parts)
full_url = urljoin(effective_base, init_url)
logger.info(f"Constructed init segment URL: {full_url}")
return full_url
logger.warning("Could not find init segment URL in manifest")
return None
@staticmethod
def extract_segment_urls(manifest_content: str, manifest_url: str) -> List[str]:
"""
DEPRECATED: Use extract_single_init_segment_url instead.
This extracts ALL segments which is inefficient.
"""
logger.warning("extract_segment_urls is deprecated, use extract_single_init_segment_url")
init_url = ManifestParser.extract_single_init_segment_url(manifest_content, manifest_url)
return [init_url] if init_url else []
@staticmethod
def extract_init_segment_urls(manifest_content: str, manifest_url: str) -> List[str]:
"""
DEPRECATED: Use extract_single_init_segment_url instead.
"""
logger.warning(
"extract_init_segment_urls is deprecated, use extract_single_init_segment_url"
)
init_url = ManifestParser.extract_single_init_segment_url(manifest_content, manifest_url)
return [init_url] if init_url else []
@@ -0,0 +1,304 @@
import base64
import struct
import uuid
from typing import List, Optional
from ..models.drm_models import PSSHData
from .logger import logger
class MP4PSSHExtractor:
"""Extract PSSH boxes and key IDs from MP4 segments"""
@staticmethod
def extract_from_url(segment_url: str, timeout: int = 10) -> List[PSSHData]:
"""Download MP4 segment and extract PSSH data"""
import requests
try:
response = requests.get(segment_url, timeout=timeout)
response.raise_for_status()
# Only download first ~100KB for efficiency
chunk_size = 1024 * 100
data = response.content[:chunk_size]
return MP4PSSHExtractor.extract_from_bytes(data)
except Exception as e:
logger.error(f"Failed to extract PSSH from {segment_url}: {e}")
return []
@staticmethod
def extract_from_bytes(data: bytes) -> List[PSSHData]:
"""Extract PSSH boxes from MP4 binary data"""
pssh_data_list = []
offset = 0
while offset < len(data):
try:
# Read box size (4 bytes, big-endian)
if offset + 4 > len(data):
break
box_size = struct.unpack(">I", data[offset : offset + 4])[0]
if box_size == 0:
box_size = len(data) - offset # Box extends to end of file
elif box_size == 1:
# Extended size (skip for now - rare in practice)
break
if offset + box_size > len(data):
break
# Read box type (4 bytes)
box_type = data[offset + 4 : offset + 8].decode("ascii", errors="ignore")
if box_type == "moov":
# Look for PSSH in moov box
moov_data = data[offset : offset + box_size]
pssh_in_moov = MP4PSSHExtractor._extract_from_moov(moov_data)
pssh_data_list.extend(pssh_in_moov)
elif box_type == "pssh":
# Found standalone PSSH box
pssh_box = MP4PSSHExtractor._parse_pssh_box(data[offset : offset + box_size])
if pssh_box:
pssh_data_list.append(pssh_box)
offset += box_size
except Exception as e:
logger.debug(f"Error parsing MP4 box at offset {offset}: {e}")
offset += 1 # Try to recover
return pssh_data_list
@staticmethod
def _extract_from_moov(moov_data: bytes) -> List[PSSHData]:
"""Extract PSSH boxes from moov container"""
pssh_list = []
offset = 8 # Skip moov header
while offset < len(moov_data):
try:
box_size = struct.unpack(">I", moov_data[offset : offset + 4])[0]
box_type = moov_data[offset + 4 : offset + 8].decode("ascii", errors="ignore")
if box_type == "trak":
# Parse track for PSSH
trak_data = moov_data[offset : offset + box_size]
pssh_in_trak = MP4PSSHExtractor._extract_from_trak(trak_data)
pssh_list.extend(pssh_in_trak)
elif box_type == "pssh":
# PSSH directly in moov
pssh_box = MP4PSSHExtractor._parse_pssh_box(
moov_data[offset : offset + box_size]
)
if pssh_box:
pssh_list.append(pssh_box)
offset += box_size
except:
break
return pssh_list
@staticmethod
def _extract_from_trak(trak_data: bytes) -> List[PSSHData]:
"""Extract PSSH from trak box"""
pssh_list = []
offset = 8
while offset < len(trak_data):
try:
box_size = struct.unpack(">I", trak_data[offset : offset + 4])[0]
box_type = trak_data[offset + 4 : offset + 8].decode("ascii", errors="ignore")
if box_type == "mdia":
mdia_data = trak_data[offset : offset + box_size]
pssh_in_mdia = MP4PSSHExtractor._extract_from_mdia(mdia_data)
pssh_list.extend(pssh_in_mdia)
offset += box_size
except:
break
return pssh_list
@staticmethod
def _extract_from_mdia(mdia_data: bytes) -> List[PSSHData]:
"""Extract PSSH from mdia box"""
pssh_list = []
offset = 8
while offset < len(mdia_data):
try:
box_size = struct.unpack(">I", mdia_data[offset : offset + 4])[0]
box_type = mdia_data[offset + 4 : offset + 8].decode("ascii", errors="ignore")
if box_type == "minf":
minf_data = mdia_data[offset : offset + box_size]
pssh_in_minf = MP4PSSHExtractor._extract_from_minf(minf_data)
pssh_list.extend(pssh_in_minf)
offset += box_size
except:
break
return pssh_list
@staticmethod
def _extract_from_minf(minf_data: bytes) -> List[PSSHData]:
"""Extract PSSH from minf box"""
pssh_list = []
offset = 8
while offset < len(minf_data):
try:
box_size = struct.unpack(">I", minf_data[offset : offset + 4])[0]
box_type = minf_data[offset + 4 : offset + 8].decode("ascii", errors="ignore")
if box_type == "stbl":
stbl_data = minf_data[offset : offset + box_size]
pssh_in_stbl = MP4PSSHExtractor._extract_from_stbl(stbl_data)
pssh_list.extend(pssh_in_stbl)
offset += box_size
except:
break
return pssh_list
@staticmethod
def _extract_from_stbl(stbl_data: bytes) -> List[PSSHData]:
"""Extract PSSH from stbl box (where protection scheme info usually is)"""
pssh_list = []
offset = 8
while offset < len(stbl_data):
try:
box_size = struct.unpack(">I", stbl_data[offset : offset + 4])[0]
box_type = stbl_data[offset + 4 : offset + 8].decode("ascii", errors="ignore")
if box_type == "sinf":
sinf_data = stbl_data[offset : offset + box_size]
pssh_in_sinf = MP4PSSHExtractor._extract_from_sinf(sinf_data)
pssh_list.extend(pssh_in_sinf)
offset += box_size
except:
break
return pssh_list
@staticmethod
def _extract_from_sinf(sinf_data: bytes) -> List[PSSHData]:
"""Extract PSSH from sinf (protection scheme information) box"""
pssh_list = []
offset = 8
while offset < len(sinf_data):
try:
box_size = struct.unpack(">I", sinf_data[offset : offset + 4])[0]
box_type = sinf_data[offset + 4 : offset + 8].decode("ascii", errors="ignore")
if box_type == "schi":
schi_data = sinf_data[offset : offset + box_size]
pssh_in_schi = MP4PSSHExtractor._extract_from_schi(schi_data)
pssh_list.extend(pssh_in_schi)
offset += box_size
except:
break
return pssh_list
@staticmethod
def _extract_from_schi(schi_data: bytes) -> List[PSSHData]:
"""Extract PSSH from schi box (where PSSH boxes are typically stored)"""
pssh_list = []
offset = 8
while offset < len(schi_data):
try:
box_size = struct.unpack(">I", schi_data[offset : offset + 4])[0]
box_type = schi_data[offset + 4 : offset + 8].decode("ascii", errors="ignore")
if box_type == "pssh":
pssh_box = MP4PSSHExtractor._parse_pssh_box(
schi_data[offset : offset + box_size]
)
if pssh_box:
pssh_list.append(pssh_box)
offset += box_size
except:
break
return pssh_list
@staticmethod
def _parse_pssh_box(pssh_bytes: bytes) -> Optional[PSSHData]:
"""Parse a PSSH box and extract system_id, pssh_box, and key_ids"""
try:
if len(pssh_bytes) < 32: # Minimum size for PSSH box
return None
# Parse box header
box_size = struct.unpack(">I", pssh_bytes[0:4])[0]
box_type = pssh_bytes[4:8].decode("ascii")
if box_type != "pssh":
return None
version = pssh_bytes[8]
flags = struct.unpack(">I", b"\x00" + pssh_bytes[9:12])[0]
# Extract system ID (bytes 12-28)
system_id_bytes = pssh_bytes[12:28]
system_id = str(uuid.UUID(bytes=system_id_bytes))
# Extract key IDs (if version > 0)
key_ids = []
current_offset = 28
if version > 0:
# Read KID count
if current_offset + 4 > len(pssh_bytes):
return None
kid_count = struct.unpack(">I", pssh_bytes[current_offset : current_offset + 4])[0]
current_offset += 4
# Read each KID
for _ in range(kid_count):
if current_offset + 16 > len(pssh_bytes):
break
kid_bytes = pssh_bytes[current_offset : current_offset + 16]
kid_uuid = str(uuid.UUID(bytes=kid_bytes))
key_ids.append(kid_uuid.replace("-", "").lower())
current_offset += 16
# Encode entire PSSH box as base64
pssh_b64 = base64.b64encode(pssh_bytes[:box_size]).decode("ascii")
return PSSHData(
system_id=system_id,
pssh_box=pssh_b64,
key_ids=key_ids,
source="mp4_segment",
)
except Exception as e:
logger.debug(f"Failed to parse PSSH box: {e}")
return None
@@ -58,9 +58,7 @@ class MPDCacheManager:
now = int(time.time())
if now >= expiry:
logger.debug(
f"Cache expired for {cache_key} (expired {now - expiry}s ago)"
)
logger.debug(f"Cache expired for {cache_key} (expired {now - expiry}s ago)")
# Clean up expired cache
self.vfs.delete(manifest_file)
self.vfs.delete(meta_file)
@@ -72,9 +70,7 @@ class MPDCacheManager:
logger.info(f"Cache hit for {cache_key} (expires in {expiry - now}s)")
return manifest_content
else:
logger.warning(
f"Cache metadata exists but manifest file missing for {cache_key}"
)
logger.warning(f"Cache metadata exists but manifest file missing for {cache_key}")
self.vfs.delete(meta_file)
return None
@@ -135,9 +131,7 @@ class MPDCacheManager:
self.vfs.delete(manifest_file)
return False
logger.info(
f"Cached MPD for {cache_key} with TTL={ttl}s (expires at {expiry})"
)
logger.info(f"Cached MPD for {cache_key} with TTL={ttl}s (expires at {expiry})")
return True
except Exception as e:
@@ -42,9 +42,7 @@ class MPDRewriter:
"""Decode base64 URL from proxy endpoint"""
return base64.urlsafe_b64decode(encoded.encode("utf-8")).decode("utf-8")
def build_proxy_url(
self, original_url: str, template_pattern: Optional[str] = None
) -> str:
def build_proxy_url(self, original_url: str, template_pattern: Optional[str] = None) -> str:
"""
Build proxy URL for an original media URL
@@ -139,9 +137,7 @@ class MPDRewriter:
if not rewritten.startswith("<?xml"):
rewritten = '<?xml version="1.0" encoding="UTF-8"?>\n' + rewritten
logger.debug(
f"Successfully rewrote MPD for provider '{self.provider_name}'"
)
logger.debug(f"Successfully rewrote MPD for provider '{self.provider_name}'")
return rewritten
except ET.ParseError as e:
@@ -177,9 +173,7 @@ class MPDRewriter:
parsed_manifest = urlparse(manifest_url)
manifest_dir = f"{parsed_manifest.scheme}://{parsed_manifest.netloc}{parsed_manifest.path.rsplit('/', 1)[0]}/"
resolved_base = urljoin(manifest_dir, base_url_text)
logger.debug(
f"Resolved relative BaseURL '{base_url_text}' to: {resolved_base}"
)
logger.debug(f"Resolved relative BaseURL '{base_url_text}' to: {resolved_base}")
return resolved_base
else:
# It's already an absolute URL
@@ -201,9 +195,7 @@ class MPDRewriter:
"""
# Find all BaseURL elements at any level
for parent in root.findall(".//*"):
for base_url_elem in list(
parent.findall("mpd:BaseURL", self.MPD_NAMESPACE)
):
for base_url_elem in list(parent.findall("mpd:BaseURL", self.MPD_NAMESPACE)):
parent.remove(base_url_elem)
logger.debug("Removed BaseURL element")
@@ -238,9 +230,7 @@ class MPDRewriter:
if "$" in resolved:
# Split into base path and template pattern
base_path, template_pattern = self.split_template_url(resolved)
element.attrib[attr] = self.build_proxy_url(
base_path, template_pattern
)
element.attrib[attr] = self.build_proxy_url(base_path, template_pattern)
logger.debug(
f"Rewrote template URL: {original_url} -> proxy with template {template_pattern}"
)
@@ -258,9 +248,7 @@ class MPDRewriter:
# SegmentURL typically doesn't have templates, but handle it just in case
if "$" in resolved:
base_path, template_pattern = self.split_template_url(resolved)
element.attrib["media"] = self.build_proxy_url(
base_path, template_pattern
)
element.attrib["media"] = self.build_proxy_url(base_path, template_pattern)
else:
element.attrib["media"] = self.build_proxy_url(resolved)
@@ -56,9 +56,7 @@ class TimestampConverter:
# Create timezone-aware datetime from epoch
if as_utc and timezone is None:
# Use UTC timezone
dt = datetime.datetime.fromtimestamp(
epoch_seconds, tz=TimestampConverter.UTC
)
dt = datetime.datetime.fromtimestamp(epoch_seconds, tz=TimestampConverter.UTC)
elif timezone is not None:
# Use specified timezone
dt = datetime.datetime.fromtimestamp(epoch_seconds, tz=timezone)
@@ -128,9 +126,7 @@ class TimestampConverter:
return dt.replace(tzinfo=TimestampConverter.UTC).timestamp()
@staticmethod
def _parse_custom_iso(
iso_string: str, format_type: Optional[str] = None
) -> datetime.datetime:
def _parse_custom_iso(iso_string: str, format_type: Optional[str] = None) -> datetime.datetime:
"""
Parse custom ISO formats not handled by fromisoformat.
@@ -172,18 +168,14 @@ class TimestampConverter:
except ValueError:
# Try without microseconds if microseconds format fails
if format_type == "microseconds":
dt = datetime.datetime.strptime(
iso_string, TimestampConverter.ISO_EXTENDED
)
dt = datetime.datetime.strptime(iso_string, TimestampConverter.ISO_EXTENDED)
else:
raise
return dt
@staticmethod
def now_iso(
format_type: str = "extended", timezone: Optional[datetime.tzinfo] = None
) -> str:
def now_iso(format_type: str = "extended", timezone: Optional[datetime.tzinfo] = None) -> str:
"""
Get current time as ISO 8601 string.
+9 -18
View File
@@ -10,6 +10,7 @@ from typing import Any, Dict, List, Optional, Tuple
# Import centralized environment manager
from .environment import get_environment_manager, is_kodi_environment
# Import centralized logger
from .logger import logger
@@ -35,9 +36,7 @@ class VFS:
self._explicit_config_dir = config_dir
self._env_manager = get_environment_manager()
logger.debug(
f"VFS initialized with config_dir={config_dir}, addon_subdir={addon_subdir}"
)
logger.debug(f"VFS initialized with config_dir={config_dir}, addon_subdir={addon_subdir}")
@property
def base_path(self) -> str:
@@ -53,9 +52,9 @@ class VFS:
if self.addon_subdir:
if is_kodi_environment():
# Kodi uses forward slashes
self._base_path = os.path.join(
profile_path, self.addon_subdir
).replace("\\", "/")
self._base_path = os.path.join(profile_path, self.addon_subdir).replace(
"\\", "/"
)
else:
# Standard filesystem
self._base_path = os.path.join(profile_path, self.addon_subdir)
@@ -208,9 +207,7 @@ class VFS:
with xbmcvfs.File(filepath, "w") as f:
bytes_written = f.write(content)
logger.debug(
f"Kodi file write: {bytes_written} bytes to {filepath}"
)
logger.debug(f"Kodi file write: {bytes_written} bytes to {filepath}")
return bytes_written > 0
else:
import pathlib
@@ -474,16 +471,12 @@ def get_global_vfs(config_dir: Optional[str] = None, addon_subdir: str = "") ->
# Convenience functions that use VFS cache
def exists(
filepath: str, config_dir: Optional[str] = None, addon_subdir: str = ""
) -> bool:
def exists(filepath: str, config_dir: Optional[str] = None, addon_subdir: str = "") -> bool:
"""Check if file exists"""
return get_vfs(config_dir, addon_subdir).exists(filepath)
def mkdirs(
dirpath: str, config_dir: Optional[str] = None, addon_subdir: str = ""
) -> bool:
def mkdirs(dirpath: str, config_dir: Optional[str] = None, addon_subdir: str = "") -> bool:
"""Create directories"""
return get_vfs(config_dir, addon_subdir).mkdirs(dirpath)
@@ -527,9 +520,7 @@ def write_json(
return get_vfs(config_dir, addon_subdir).write_json(filepath, data, indent)
def delete(
filepath: str, config_dir: Optional[str] = None, addon_subdir: str = ""
) -> bool:
def delete(filepath: str, config_dir: Optional[str] = None, addon_subdir: str = "") -> bool:
"""Delete file"""
return get_vfs(config_dir, addon_subdir).delete(filepath)
@@ -20,14 +20,10 @@ class AuthStatusBuilder:
current_auth_type = provider.get_current_auth_type(context)
# Calculate auth state (with provider override support)
auth_state = AuthStatusBuilder._calculate_auth_state_with_override(
provider, context
)
auth_state = AuthStatusBuilder._calculate_auth_state_with_override(provider, context)
# Calculate readiness (with provider override support)
is_ready, reason = AuthStatusBuilder._calculate_readiness_with_override(
provider, context
)
is_ready, reason = AuthStatusBuilder._calculate_readiness_with_override(provider, context)
# Build token info
token_scopes = AuthStatusBuilder._build_token_scopes(provider, context)
@@ -65,9 +61,7 @@ class AuthStatusBuilder:
has_valid_token=AuthStatusBuilder._has_valid_token(provider, context),
primary_token_scope=provider.primary_token_scope,
token_scopes=token_scopes,
last_authentication=AuthStatusBuilder._get_last_auth_time(
provider, context
),
last_authentication=AuthStatusBuilder._get_last_auth_time(provider, context),
provider_specific=provider_specific,
token_expires_at=token_expires_at,
token_expires_in_seconds=token_expires_in_seconds,
@@ -78,9 +72,7 @@ class AuthStatusBuilder:
# ===== Calculation Methods (with provider override support) =====
@staticmethod
def _calculate_auth_state_with_override(
provider, context: AuthContext
) -> AuthState:
def _calculate_auth_state_with_override(provider, context: AuthContext) -> AuthState:
"""Calculate auth state, allowing provider override"""
# Check if provider has custom logic
if hasattr(provider, "_calculate_auth_state"):
@@ -92,9 +84,7 @@ class AuthStatusBuilder:
return AuthStatusBuilder._calculate_auth_state(provider, context)
@staticmethod
def _calculate_readiness_with_override(
provider, context: AuthContext
) -> Tuple[bool, str]:
def _calculate_readiness_with_override(provider, context: AuthContext) -> Tuple[bool, str]:
"""Calculate readiness, allowing provider override"""
# Check if provider has custom logic
if hasattr(provider, "_calculate_readiness"):
@@ -120,9 +110,7 @@ class AuthStatusBuilder:
return AuthState.AUTHENTICATED
# Check if we have an expired token that can be refreshed
expired_token = AuthStatusBuilder._get_expired_token_with_refresh(
provider, context
)
expired_token = AuthStatusBuilder._get_expired_token_with_refresh(provider, context)
if expired_token:
return AuthState.EXPIRED
@@ -151,17 +139,13 @@ class AuthStatusBuilder:
return True, "Has valid authentication token"
# 3. Check if we have an expired token that can be refreshed
expired_token = AuthStatusBuilder._get_expired_token_with_refresh(
provider, context
)
expired_token = AuthStatusBuilder._get_expired_token_with_refresh(provider, context)
if expired_token:
return True, "Has expired token with refresh capability"
# 4. Check credentials if required (only matters if no valid token)
if provider.requires_stored_credentials:
credentials = context.get_credentials(
provider.provider_name, provider.country
)
credentials = context.get_credentials(provider.provider_name, provider.country)
if not credentials:
return False, "Missing required credentials"
@@ -202,9 +186,7 @@ class AuthStatusBuilder:
return False
@staticmethod
def _get_expired_token_with_refresh(
provider, context: AuthContext
) -> Optional[Dict[str, Any]]:
def _get_expired_token_with_refresh(provider, context: AuthContext) -> Optional[Dict[str, Any]]:
"""Get an expired token that has refresh capability"""
# Check primary scope
if provider.primary_token_scope:
@@ -233,9 +215,7 @@ class AuthStatusBuilder:
token_scopes = {}
for scope in provider.token_scopes:
token_data = context.get_token(
provider.provider_name, scope, provider.country
)
token_data = context.get_token(provider.provider_name, scope, provider.country)
if token_data:
expires_at = None
if "issued_at" in token_data and "expires_in" in token_data:
@@ -256,9 +236,7 @@ class AuthStatusBuilder:
return token_scopes
@staticmethod
def _get_credentials_info(
provider, context: AuthContext
) -> Tuple[bool, Optional[str]]:
def _get_credentials_info(provider, context: AuthContext) -> Tuple[bool, Optional[str]]:
"""Get credentials information"""
if not provider.requires_stored_credentials:
return False, None
@@ -335,16 +313,12 @@ class AuthStatusBuilder:
# If no primary scope or no token, check root-level token
if not token_data:
token_data = context.get_token(
provider.provider_name, None, provider.country
)
token_data = context.get_token(provider.provider_name, None, provider.country)
# If still no token, try first available scope
if not token_data and provider.token_scopes:
for scope in provider.token_scopes:
token_data = context.get_token(
provider.provider_name, scope, provider.country
)
token_data = context.get_token(provider.provider_name, scope, provider.country)
if token_data:
break
@@ -356,22 +330,15 @@ class AuthStatusBuilder:
if "expires_in" in token_data and "issued_at" in token_data:
expires_in = token_data["expires_in"]
issued_at = token_data["issued_at"]
if isinstance(expires_in, (int, float)) and isinstance(
issued_at, (int, float)
):
if isinstance(expires_in, (int, float)) and isinstance(issued_at, (int, float)):
token_expires_at = float(issued_at) + float(expires_in)
token_expires_in_seconds = int(token_expires_at - current_time)
# yo_digital format: separate access token expiration
elif (
"access_token_expires_in" in token_data
and "access_token_issued_at" in token_data
):
elif "access_token_expires_in" in token_data and "access_token_issued_at" in token_data:
expires_in = token_data["access_token_expires_in"]
issued_at = token_data["access_token_issued_at"]
if isinstance(expires_in, (int, float)) and isinstance(
issued_at, (int, float)
):
if isinstance(expires_in, (int, float)) and isinstance(issued_at, (int, float)):
token_expires_at = float(issued_at) + float(expires_in)
token_expires_in_seconds = int(token_expires_at - current_time)
@@ -383,21 +350,14 @@ class AuthStatusBuilder:
token_expires_in_seconds = int(token_expires_at - current_time)
# Calculate refresh token expiration (yo_digital format)
if (
"refresh_token_expires_in" in token_data
and "refresh_token_issued_at" in token_data
):
if "refresh_token_expires_in" in token_data and "refresh_token_issued_at" in token_data:
refresh_expires_in = token_data["refresh_token_expires_in"]
refresh_issued_at = token_data["refresh_token_issued_at"]
if isinstance(refresh_expires_in, (int, float)) and isinstance(
refresh_issued_at, (int, float)
):
refresh_token_expires_at = float(refresh_issued_at) + float(
refresh_expires_in
)
refresh_token_expires_in_seconds = int(
refresh_token_expires_at - current_time
)
refresh_token_expires_at = float(refresh_issued_at) + float(refresh_expires_in)
refresh_token_expires_in_seconds = int(refresh_token_expires_at - current_time)
return (
token_expires_at,
@@ -11,9 +11,7 @@ class AuthContext:
def __init__(self, settings_manager):
self.settings = settings_manager
self.session = settings_manager.session_manager if settings_manager else None
self.credentials = (
settings_manager.credential_manager if settings_manager else None
)
self.credentials = settings_manager.credential_manager if settings_manager else None
def get_credentials(self, provider_name: str, country: str = None) -> Optional[Any]:
"""Get credentials for provider"""
+34 -102
View File
@@ -8,8 +8,7 @@ import uuid
from datetime import datetime, timedelta
from typing import Any, Dict, Optional
from ...base.auth.base_auth import (BaseAuthenticator, BaseAuthToken,
TokenAuthLevel)
from ...base.auth.base_auth import BaseAuthenticator, BaseAuthToken, TokenAuthLevel
from ...base.models.proxy_models import ProxyConfig
from ...base.utils.logger import logger
from .constants import HRTiConfig
@@ -108,9 +107,7 @@ class HRTiAuthenticator(BaseAuthenticator):
return headers
def _get_api_headers(
self, bearer_token: str = None, referer: str = None
) -> Dict[str, str]:
def _get_api_headers(self, bearer_token: str = None, referer: str = None) -> Dict[str, str]:
"""Get headers for regular API endpoints with proper session support"""
if not self._ip_address:
self._get_ip_address()
@@ -139,18 +136,14 @@ class HRTiAuthenticator(BaseAuthenticator):
# Add authorization header with Client prefix
if bearer_token:
headers["authorization"] = f"Client {bearer_token}"
logger.debug(
f"Added authorization header with token: {bearer_token[:20]}..."
)
logger.debug(f"Added authorization header with token: {bearer_token[:20]}...")
return headers
def _build_auth_payload(self) -> Dict[str, Any]:
"""Build HRTi-specific authentication payload"""
# HRTi uses specific format: {"Username":"...","Password":"...","OperatorReferenceId":"hrt"}
if hasattr(self.credentials, "username") and hasattr(
self.credentials, "password"
):
if hasattr(self.credentials, "username") and hasattr(self.credentials, "password"):
return {
"Username": self.credentials.username,
"Password": self.credentials.password,
@@ -164,9 +157,7 @@ class HRTiAuthenticator(BaseAuthenticator):
"OperatorReferenceId": self.config.operator_reference_id,
}
def _create_token_from_response(
self, response_data: Dict[str, Any]
) -> HRTiAuthToken:
def _create_token_from_response(self, response_data: Dict[str, Any]) -> HRTiAuthToken:
"""Create HRTi-specific token from API response"""
result = response_data.get("Result", {})
@@ -175,13 +166,9 @@ class HRTiAuthenticator(BaseAuthenticator):
# Extract token
access_token = result.get("Token", "")
logger.debug(
f"Extracted access_token: {'present' if access_token else 'MISSING'}"
)
logger.debug(f"Extracted access_token: {'present' if access_token else 'MISSING'}")
if access_token:
logger.debug(
f"Token length: {len(access_token)}, starts with: {access_token[:20]}..."
)
logger.debug(f"Token length: {len(access_token)}, starts with: {access_token[:20]}...")
# Store user ID if available
if "Customer" in result:
@@ -231,9 +218,7 @@ class HRTiAuthenticator(BaseAuthenticator):
self._current_token = token # THIS IS THE CRITICAL MISSING LINE
# Debug: Verify token is available
logger.debug(
f"Token created - access_token present: {bool(token.access_token)}"
)
logger.debug(f"Token created - access_token present: {bool(token.access_token)}")
if token.access_token:
logger.debug(f"Token value: {token.access_token[:20]}...")
else:
@@ -264,9 +249,7 @@ class HRTiAuthenticator(BaseAuthenticator):
configured_ip = self._get_configured_ip_from_settings()
if configured_ip:
self._ip_address = configured_ip
logger.info(
f"Using configured IP address from Kodi settings: {self._ip_address}"
)
logger.info(f"Using configured IP address from Kodi settings: {self._ip_address}")
return self._ip_address
except Exception as e:
logger.debug(f"Could not read configured IP from settings: {e}")
@@ -274,13 +257,9 @@ class HRTiAuthenticator(BaseAuthenticator):
# Priority 2: Fall back to API fetch if no configured IP
try:
logger.debug("No configured IP found, fetching from API...")
response = self.http_manager.get(
self.config.api_endpoints["get_ip"], operation="api"
)
response = self.http_manager.get(self.config.api_endpoints["get_ip"], operation="api")
response.raise_for_status()
self._ip_address = response.text.strip().strip(
'"'
) # Remove quotes if present
self._ip_address = response.text.strip().strip('"') # Remove quotes if present
logger.info(f"Retrieved IP address from API: {self._ip_address}")
return self._ip_address
except Exception as e:
@@ -319,16 +298,12 @@ class HRTiAuthenticator(BaseAuthenticator):
"""Get HRTi environment configuration"""
try:
# Get env config
env_response = self.http_manager.get(
self.config.env_endpoint, operation="api"
)
env_response = self.http_manager.get(self.config.env_endpoint, operation="api")
env_response.raise_for_status()
env_data = env_response.json()
# Get main config
config_response = self.http_manager.get(
self.config.config_endpoint, operation="api"
)
config_response = self.http_manager.get(self.config.config_endpoint, operation="api")
config_response.raise_for_status()
config_data = config_response.json()
@@ -348,9 +323,7 @@ class HRTiAuthenticator(BaseAuthenticator):
device_id = self.get_device_id()
logger.debug(
f"Performing grant access with username: {self.credentials.username}"
)
logger.debug(f"Performing grant access with username: {self.credentials.username}")
logger.debug(f"Using device ID: {device_id}")
logger.debug(f"Using IP address: {self._ip_address}")
@@ -366,9 +339,7 @@ class HRTiAuthenticator(BaseAuthenticator):
safe_payload = payload.copy()
if "Password" in safe_payload:
safe_payload["Password"] = (
"***" if safe_payload["Password"] else "<empty>"
)
safe_payload["Password"] = "***" if safe_payload["Password"] else "<empty>"
logger.debug(f"HRTi Auth Payload: {safe_payload}")
response = self.http_manager.post(
@@ -385,11 +356,7 @@ class HRTiAuthenticator(BaseAuthenticator):
if "Result" not in result:
raise Exception("No result in grant access response")
customer_id = (
result.get("Result", {})
.get("Customer", {})
.get("CustomerId", "unknown")
)
customer_id = result.get("Result", {}).get("Customer", {}).get("CustomerId", "unknown")
logger.debug(f"Grant access successful - user: {customer_id}")
return result
@@ -409,9 +376,7 @@ class HRTiAuthenticator(BaseAuthenticator):
def _register_device(self):
"""Register device with HRTi using API headers with proper authorization"""
try:
bearer_token = (
self._current_token.access_token if self._current_token else ""
)
bearer_token = self._current_token.access_token if self._current_token else ""
if not bearer_token:
logger.warning("No bearer token available for device registration")
return
@@ -458,9 +423,7 @@ class HRTiAuthenticator(BaseAuthenticator):
def _get_initial_data(self):
"""Get initial content rating and profiles using API headers with proper authorization"""
try:
bearer_token = (
self._current_token.access_token if self._current_token else ""
)
bearer_token = self._current_token.access_token if self._current_token else ""
headers = self._get_api_headers(bearer_token)
# Update the referer to root path for these API calls
@@ -492,18 +455,14 @@ class HRTiAuthenticator(BaseAuthenticator):
def _initialize_device(self):
"""Initialize or load device ID - ensure it's a proper UUID"""
try:
self._device_id = self.settings_manager.get_device_id(
self.provider_name, self.country
)
self._device_id = self.settings_manager.get_device_id(self.provider_name, self.country)
if not self._device_id:
self._device_id = str(uuid.uuid4())
logger.debug(f"Generated new device ID: {self._device_id}")
# Ensure it's a valid UUID format
if not self._validate_uuid(self._device_id):
logger.warning(
f"Invalid device ID format, generating new one: {self._device_id}"
)
logger.warning(f"Invalid device ID format, generating new one: {self._device_id}")
self._device_id = str(uuid.uuid4())
except Exception as e:
@@ -557,9 +516,7 @@ class HRTiAuthenticator(BaseAuthenticator):
) -> Optional[Dict[str, Any]]:
"""Authorize a playback session"""
try:
bearer_token = (
self._current_token.access_token if self._current_token else ""
)
bearer_token = self._current_token.access_token if self._current_token else ""
headers = self._get_api_headers(bearer_token)
# Set referer based on content type
@@ -567,9 +524,7 @@ class HRTiAuthenticator(BaseAuthenticator):
headers["referer"] = f"{self.config.base_website}/videostore"
else:
if content_type == "tlive":
headers["referer"] = (
f"{self.config.base_website}/live/tv?channel={channel_id}"
)
headers["referer"] = f"{self.config.base_website}/live/tv?channel={channel_id}"
elif content_type == "rlive":
headers["referer"] = f"{self.config.base_website}/live/radio"
else:
@@ -600,9 +555,7 @@ class HRTiAuthenticator(BaseAuthenticator):
result = response.json()
if "Result" in result:
authorized = result["Result"].get("Authorized", False)
session_id = result["Result"].get("SessionId") or result["Result"].get(
"DrmId"
)
session_id = result["Result"].get("SessionId") or result["Result"].get("DrmId")
logger.debug(
f"Session authorization result - authorized: {authorized}, session_id: {session_id}"
)
@@ -618,18 +571,14 @@ class HRTiAuthenticator(BaseAuthenticator):
def report_session_event(self, session_id: str, channel_id: str = None) -> bool:
"""Report session event (like play start)"""
try:
bearer_token = (
self._current_token.access_token if self._current_token else ""
)
bearer_token = self._current_token.access_token if self._current_token else ""
headers = self._get_api_headers(bearer_token)
# Set referer based on channel
if channel_id is None:
headers["referer"] = f"{self.config.base_website}/videostore"
else:
headers["referer"] = (
f"{self.config.base_website}/live/tv?channel={channel_id}"
)
headers["referer"] = f"{self.config.base_website}/live/tv?channel={channel_id}"
payload = {"SessionEventId": 1, "SessionId": session_id} # 1 = play start
@@ -665,9 +614,7 @@ class HRTiAuthenticator(BaseAuthenticator):
license_bytes = json.dumps(drm_license).encode("utf-8")
license_b64 = base64.b64encode(license_bytes).decode("utf-8")
logger.debug(
f"License data (base64, first 30 chars): {license_b64[:30]}..."
)
logger.debug(f"License data (base64, first 30 chars): {license_b64[:30]}...")
return license_b64
except Exception as e:
@@ -700,13 +647,10 @@ class HRTiAuthenticator(BaseAuthenticator):
self.provider_name, self.country
)
except TypeError:
stored_creds = self.settings_manager.get_provider_credentials(
self.provider_name
)
stored_creds = self.settings_manager.get_provider_credentials(self.provider_name)
has_user_creds = (
isinstance(stored_creds, UserPasswordCredentials)
and stored_creds.validate()
isinstance(stored_creds, UserPasswordCredentials) and stored_creds.validate()
)
# Check current credentials
@@ -721,10 +665,7 @@ class HRTiAuthenticator(BaseAuthenticator):
return TokenAuthLevel.USER_AUTHENTICATED
# Check if using anonymous credentials
if (
hasattr(self.credentials, "username")
and self.credentials.username == "anonymoushrt"
):
if hasattr(self.credentials, "username") and self.credentials.username == "anonymoushrt":
return TokenAuthLevel.ANONYMOUS
# If we have user credentials, consider it user authenticated
@@ -758,13 +699,10 @@ class HRTiAuthenticator(BaseAuthenticator):
self.provider_name, self.country
)
except TypeError:
stored_creds = self.settings_manager.get_provider_credentials(
self.provider_name
)
stored_creds = self.settings_manager.get_provider_credentials(self.provider_name)
has_stored_user_creds = (
isinstance(stored_creds, UserPasswordCredentials)
and stored_creds.validate()
isinstance(stored_creds, UserPasswordCredentials) and stored_creds.validate()
)
# Check current credentials
@@ -778,9 +716,7 @@ class HRTiAuthenticator(BaseAuthenticator):
return False
def get_bearer_token(
self, force_refresh: bool = False, force_upgrade: bool = False
) -> str:
def get_bearer_token(self, force_refresh: bool = False, force_upgrade: bool = False) -> str:
"""
Get bearer token with automatic upgrade from anonymous to user credentials
"""
@@ -809,9 +745,7 @@ class HRTiAuthenticator(BaseAuthenticator):
self.provider_name, self.country
)
except TypeError:
user_creds = self.settings_manager.get_provider_credentials(
self.provider_name
)
user_creds = self.settings_manager.get_provider_credentials(self.provider_name)
if not user_creds or not user_creds.validate():
logger.debug("No valid user credentials available for upgrade")
@@ -821,9 +755,7 @@ class HRTiAuthenticator(BaseAuthenticator):
self.credentials = user_creds
# Perform authentication with user credentials
logger.info(
"Performing authentication with user credentials for upgrade"
)
logger.info("Performing authentication with user credentials for upgrade")
user_token = self._perform_authentication()
if user_token and not user_token.is_expired:
@@ -68,14 +68,10 @@ class HRTiConfig:
self.base_url = config.get("base_url", HRTiDefaults.BASE_URL)
self.hsapi_base_url = config.get("hsapi_base_url", HRTiDefaults.HSAPI_BASE_URL)
self.env_endpoint = config.get("env_endpoint", HRTiDefaults.ENV_ENDPOINT)
self.config_endpoint = config.get(
"config_endpoint", HRTiDefaults.CONFIG_ENDPOINT
)
self.config_endpoint = config.get("config_endpoint", HRTiDefaults.CONFIG_ENDPOINT)
# API endpoints configuration
self.api_endpoints = config.get(
"api_endpoints", HRTiDefaults.API_ENDPOINTS.copy()
)
self.api_endpoints = config.get("api_endpoints", HRTiDefaults.API_ENDPOINTS.copy())
# DRM and License
self.license_url = config.get("license_url", HRTiDefaults.LICENSE_URL)
@@ -88,9 +84,7 @@ class HRTiConfig:
"operator_reference_id", HRTiDefaults.OPERATOR_REFERENCE_ID
)
self.merchant = config.get("merchant", HRTiDefaults.MERCHANT)
self.connection_type = config.get(
"connection_type", HRTiDefaults.CONNECTION_TYPE
)
self.connection_type = config.get("connection_type", HRTiDefaults.CONNECTION_TYPE)
self.application_version = config.get(
"application_version", HRTiDefaults.APPLICATION_VERSION
)
@@ -4,8 +4,7 @@ from typing import ClassVar, Dict, List, Optional
import requests
from ...base.models import (DRMConfig, DRMSystem, LicenseConfig,
LicenseUnwrapperParams)
from ...base.models import DRMConfig, DRMSystem, LicenseConfig, LicenseUnwrapperParams
from ...base.models.proxy_models import ProxyConfig
from ...base.models.streaming_channel import StreamingChannel
from ...base.provider import AuthType, StreamingProvider
@@ -56,9 +55,7 @@ class HRTiProvider(StreamingProvider):
)
# Share HTTP manager for consistency
self.http_manager = self._share_http_manager_with_authenticator(
self.authenticator
)
self.http_manager = self._share_http_manager_with_authenticator(self.authenticator)
try:
# Initialize authentication
@@ -176,9 +173,7 @@ class HRTiProvider(StreamingProvider):
channels.append(streaming_channel)
self.channels = channels
logger.info(
f"Successfully fetched {len(channels)} channels from HRTi on retry"
)
logger.info(f"Successfully fetched {len(channels)} channels from HRTi on retry")
return channels
else:
logger.warning("No channels found in HRTi retry response")
@@ -228,9 +223,7 @@ class HRTiProvider(StreamingProvider):
channel.use_cdm = True
channel.cdm_type = "widevine"
logger.debug(
f"Parsed HRTi channel: {name} ({channel_id}) - radio: {is_radio}"
)
logger.debug(f"Parsed HRTi channel: {name} ({channel_id}) - radio: {is_radio}")
return channel
except Exception as e:
@@ -276,9 +269,7 @@ class HRTiProvider(StreamingProvider):
)
if not session_data:
logger.warning(
f"Failed to authorize session for channel {channel.name}"
)
logger.warning(f"Failed to authorize session for channel {channel.name}")
return channel
# Check if authorized
@@ -301,9 +292,7 @@ class HRTiProvider(StreamingProvider):
# Set DRM configuration with session data
# Pass session_data to avoid re-authorizing
drm_configs = self.get_drm(
channel.channel_id, session_data=session_data, **kwargs
)
drm_configs = self.get_drm(channel.channel_id, session_data=session_data, **kwargs)
if drm_configs:
channel.use_cdm = True
channel.cdm_type = "widevine"
@@ -367,9 +356,7 @@ class HRTiProvider(StreamingProvider):
logger.error(f"Error getting manifest for channel {channel_id}: {e}")
return None
def get_drm(
self, channel_id: str, session_data: Dict = None, **kwargs
) -> List[DRMConfig]:
def get_drm(self, channel_id: str, session_data: Dict = None, **kwargs) -> List[DRMConfig]:
"""
Get DRM configurations for a channel with proper license data.
If session_data is not provided, will authorize a new session.
@@ -394,15 +381,11 @@ class HRTiProvider(StreamingProvider):
break
if not target_channel:
logger.error(
f"Channel {channel_id} not found for DRM authorization"
)
logger.error(f"Channel {channel_id} not found for DRM authorization")
return []
# Determine content type
content_type = (
"rlive" if target_channel.content_type == "AUDIO" else "tlive"
)
content_type = "rlive" if target_channel.content_type == "AUDIO" else "tlive"
# Parse the streaming URL to get content DRM ID
from urllib.parse import urlparse
@@ -431,16 +414,12 @@ class HRTiProvider(StreamingProvider):
)
if not session_data:
logger.error(
f"Failed to authorize session for DRM - channel {channel_id}"
)
logger.error(f"Failed to authorize session for DRM - channel {channel_id}")
return []
# Check if authorized
if not session_data.get("Authorized", False):
logger.warning(
f"Session not authorized for DRM - channel {channel_id}"
)
logger.warning(f"Session not authorized for DRM - channel {channel_id}")
return []
logger.debug(f"Session authorized for DRM - channel {channel_id}")
@@ -497,13 +476,9 @@ class HRTiProvider(StreamingProvider):
)
# Create the DRM configuration
drm_config = DRMConfig(
system=DRMSystem.WIDEVINE, priority=1, license=license_config
)
drm_config = DRMConfig(system=DRMSystem.WIDEVINE, priority=1, license=license_config)
logger.debug(
f"Created DRM config for channel {channel_id} with DrmId {drm_id}"
)
logger.debug(f"Created DRM config for channel {channel_id} with DrmId {drm_id}")
return [drm_config]
except Exception as e:
@@ -592,9 +567,7 @@ class HRTiProvider(StreamingProvider):
# HRTi doesn't provide XMLTV format natively
return None
def get_dynamic_manifest_params(
self, channel: StreamingChannel, **kwargs
) -> Optional[str]:
def get_dynamic_manifest_params(self, channel: StreamingChannel, **kwargs) -> Optional[str]:
"""
Get dynamic manifest parameters for HRTi channels.
@@ -1,7 +1,11 @@
# streaming_providers/providers/joyn/__init__.py
from .auth import JoynAuthenticator, JoynAuthToken, JoynCredentials
from .constants import (COUNTRY_TENANT_MAPPING, DEFAULT_VIDEO_CONFIG,
JOYN_GRAPHQL_ENDPOINTS, JOYN_STREAMING_ENDPOINTS)
from .constants import (
COUNTRY_TENANT_MAPPING,
DEFAULT_VIDEO_CONFIG,
JOYN_GRAPHQL_ENDPOINTS,
JOYN_STREAMING_ENDPOINTS,
)
from .models import JoynChannel, PlaybackRestrictedException
from .provider import JoynProvider
+37 -73
View File
@@ -12,11 +12,20 @@ from ...base.auth.base_oauth2_auth import BaseOAuth2Authenticator
from ...base.auth.credentials import ClientCredentials
from ...base.models.proxy_models import ProxyConfig
from ...base.utils.logger import logger
from .constants import (COUNTRY_TENANT_MAPPING, DEFAULT_COUNTRY,
DEFAULT_PLATFORM, DEFAULT_REQUEST_TIMEOUT, DEVICE_IDS,
JOYN_AUTH_ENDPOINTS, JOYN_CLIENT_VERSION, JOYN_DOMAINS,
JOYN_OAUTH_SCOPE, JOYN_SSO_DISCOVERY_URL,
JOYN_USER_AGENT, SUPPORTED_COUNTRIES)
from .constants import (
COUNTRY_TENANT_MAPPING,
DEFAULT_COUNTRY,
DEFAULT_PLATFORM,
DEFAULT_REQUEST_TIMEOUT,
DEVICE_IDS,
JOYN_AUTH_ENDPOINTS,
JOYN_CLIENT_VERSION,
JOYN_DOMAINS,
JOYN_OAUTH_SCOPE,
JOYN_SSO_DISCOVERY_URL,
JOYN_USER_AGENT,
SUPPORTED_COUNTRIES,
)
class JoynSSODiscovery:
@@ -110,9 +119,7 @@ class JoynCredentials(ClientCredentials):
def __post_init__(self):
# Set client_id from constant if not provided
if not self.client_id:
self.client_id = DEVICE_IDS.get(
self.client_name, DEVICE_IDS[DEFAULT_PLATFORM]
)
self.client_id = DEVICE_IDS.get(self.client_name, DEVICE_IDS[DEFAULT_PLATFORM])
if not self.distribution_tenant and self.country in COUNTRY_TENANT_MAPPING:
self.distribution_tenant = COUNTRY_TENANT_MAPPING[self.country]
@@ -349,9 +356,7 @@ class JoynAuthenticator(BaseOAuth2Authenticator):
if not client_id:
raise Exception("No client_id found in login endpoint")
logger.debug(
f"Extracted and cached client_id for {self.platform}: {client_id}"
)
logger.debug(f"Extracted and cached client_id for {self.platform}: {client_id}")
return client_id
except Exception as e:
@@ -399,9 +404,7 @@ class JoynAuthenticator(BaseOAuth2Authenticator):
raise Exception("No credentials available")
return self.credentials.to_auth_payload()
def _create_token_from_response(
self, response_data: Dict[str, Any]
) -> BaseAuthToken:
def _create_token_from_response(self, response_data: Dict[str, Any]) -> BaseAuthToken:
"""Create token object from API response"""
token = JoynAuthToken(
access_token=response_data["access_token"],
@@ -434,22 +437,17 @@ class JoynAuthenticator(BaseOAuth2Authenticator):
"""
Perform authentication using appropriate flow based on credential type
"""
from ...base.auth.credentials import (ClientCredentials,
UserPasswordCredentials)
from ...base.auth.credentials import ClientCredentials, UserPasswordCredentials
if isinstance(self.credentials, UserPasswordCredentials):
# Use OAuth2 authorization code flow with PKCE
logger.info(
f"Using OAuth2 authorization code flow for {self.provider_name}"
)
logger.info(f"Using OAuth2 authorization code flow for {self.provider_name}")
token_data = self._perform_oauth_authorization_code_flow(
self.credentials.username, self.credentials.password
)
elif isinstance(self.credentials, ClientCredentials):
# Use client credentials flow (anonymous auth)
logger.info(
f"Using OAuth2 client credentials flow for {self.provider_name}"
)
logger.info(f"Using OAuth2 client credentials flow for {self.provider_name}")
token_data = self._perform_oauth_client_credentials_flow()
else:
raise Exception(
@@ -483,9 +481,7 @@ class JoynAuthenticator(BaseOAuth2Authenticator):
response.raise_for_status()
token_data = response.json()
logger.debug(
f"OAuth2 client credentials flow successful for {self.provider_name}"
)
logger.debug(f"OAuth2 client credentials flow successful for {self.provider_name}")
return token_data
except Exception as e:
@@ -534,18 +530,14 @@ class JoynAuthenticator(BaseOAuth2Authenticator):
initiate_result = initiate_response.json()
if not initiate_result.get("success"):
raise Exception(
f"Password verification initiation failed: {initiate_result}"
)
raise Exception(f"Password verification initiation failed: {initiate_result}")
exchange_data = initiate_result["data"]
exchange_id = exchange_data["exchange_id"]["exchange_id"]
sub = exchange_data["sub"]
status_id = exchange_data["status_id"]
logger.debug(
f"Got exchange_id: {exchange_id}, sub: {sub}, status_id: {status_id}"
)
logger.debug(f"Got exchange_id: {exchange_id}, sub: {sub}, status_id: {status_id}")
# Step 2: Authenticate with password
authenticate_data = {
@@ -564,9 +556,7 @@ class JoynAuthenticator(BaseOAuth2Authenticator):
timeout=self._config.timeout,
)
logger.debug(
f"Authenticate response status: {authenticate_response.status_code}"
)
logger.debug(f"Authenticate response status: {authenticate_response.status_code}")
authenticate_response.raise_for_status()
authenticate_result = authenticate_response.json()
@@ -656,9 +646,7 @@ class JoynAuthenticator(BaseOAuth2Authenticator):
fragment_params = parse_qs(fragment)
auth_code = fragment_params.get("code", [None])[0]
if auth_code:
logger.debug(
f"Found authorization code in fragment: {auth_code}"
)
logger.debug(f"Found authorization code in fragment: {auth_code}")
return auth_code
# If no code found, try to follow redirects
@@ -712,18 +700,14 @@ class JoynAuthenticator(BaseOAuth2Authenticator):
logger.debug("Fetching authorization page")
auth_response = session.get(authorization_url, timeout=self._config.timeout)
logger.debug(
f"Authorization page response status: {auth_response.status_code}"
)
logger.debug(f"Authorization page response status: {auth_response.status_code}")
logger.debug(f"Authorization page response URL: {auth_response.url}")
auth_response.raise_for_status()
# Step 2a: Check if we got redirected directly to callback (existing session)
if self.oauth_redirect_uri in auth_response.url:
logger.info(
"User already authenticated - extracting code from redirect"
)
logger.info("User already authenticated - extracting code from redirect")
# Parse the redirect URL for authorization code
parsed_url = urlparse(auth_response.url)
@@ -733,16 +717,12 @@ class JoynAuthenticator(BaseOAuth2Authenticator):
received_state = query_params.get("state", [None])[0]
if not auth_code:
raise Exception(
"Redirect to callback but no authorization code found"
)
raise Exception("Redirect to callback but no authorization code found")
if not self.validate_oauth_state(received_state, state):
raise Exception("State validation failed on direct redirect")
logger.debug(
f"Extracted authorization code from direct redirect: {auth_code}"
)
logger.debug(f"Extracted authorization code from direct redirect: {auth_code}")
# Exchange code for token
logger.debug("Exchanging authorization code for tokens")
@@ -752,9 +732,7 @@ class JoynAuthenticator(BaseOAuth2Authenticator):
state=state,
)
logger.debug(
"Joyn OAuth2 authorization code flow successful (existing session)"
)
logger.debug("Joyn OAuth2 authorization code flow successful (existing session)")
return token_data
# Step 2b: No existing session - need to perform login
@@ -773,9 +751,7 @@ class JoynAuthenticator(BaseOAuth2Authenticator):
if request_id_match:
request_id = request_id_match.group(1)
else:
raise Exception(
"Could not extract request_id from authorization page"
)
raise Exception("Could not extract request_id from authorization page")
logger.debug(f"Extracted request_id: {request_id}")
@@ -853,9 +829,7 @@ class JoynAuthenticator(BaseOAuth2Authenticator):
# Use DEVICE_IDS directly instead of oauth_client_id
# Joyn refresh might need the platform-specific device ID, not the OAuth client ID
payload = {
"client_id": DEVICE_IDS.get(
self.platform, DEVICE_IDS[DEFAULT_PLATFORM]
),
"client_id": DEVICE_IDS.get(self.platform, DEVICE_IDS[DEFAULT_PLATFORM]),
"client_name": self.platform,
"grant_type": "Bearer", # Joyn uses 'Bearer' instead of 'refresh_token'
"refresh_token": self._current_token.refresh_token,
@@ -965,9 +939,7 @@ class JoynAuthenticator(BaseOAuth2Authenticator):
# 2. Check for social_id presence - clear indicator of user authentication
if "social_id" in claims:
logger.debug(
"Token classified as USER_AUTHENTICATED (social_id present)"
)
logger.debug("Token classified as USER_AUTHENTICATED (social_id present)")
return TokenAuthLevel.USER_AUTHENTICATED
# 3. Check client ID (cId) against known client IDs
@@ -987,9 +959,7 @@ class JoynAuthenticator(BaseOAuth2Authenticator):
if subject and len(subject) == 36: # UUID format
# Client credentials tokens often have UUID subjects representing the client
# User tokens might have different patterns or include user identifiers
logger.debug(
"Token classified as CLIENT_CREDENTIALS (UUID subject pattern)"
)
logger.debug("Token classified as CLIENT_CREDENTIALS (UUID subject pattern)")
return TokenAuthLevel.CLIENT_CREDENTIALS
# 5. Fallback: Check token scope or other claims
@@ -997,14 +967,10 @@ class JoynAuthenticator(BaseOAuth2Authenticator):
if scope:
scopes = scope.split()
if "offline_access" in scopes and "profile" in scopes:
logger.debug(
"Token classified as USER_AUTHENTICATED (user scopes present)"
)
logger.debug("Token classified as USER_AUTHENTICATED (user scopes present)")
return TokenAuthLevel.USER_AUTHENTICATED
elif "openid" in scopes and len(scopes) <= 2:
logger.debug(
"Token classified as CLIENT_CREDENTIALS (minimal scopes)"
)
logger.debug("Token classified as CLIENT_CREDENTIALS (minimal scopes)")
return TokenAuthLevel.CLIENT_CREDENTIALS
logger.warning(f"Could not definitively classify token, using UNKNOWN")
@@ -1037,9 +1003,7 @@ class JoynAuthenticator(BaseOAuth2Authenticator):
"cId": claims.get("cId", "MISSING"),
"social_id": "PRESENT" if "social_id" in claims else "MISSING",
"sub": (
claims.get("sub", "MISSING")[:8] + "..."
if claims.get("sub")
else "MISSING"
claims.get("sub", "MISSING")[:8] + "..." if claims.get("sub") else "MISSING"
),
"scope": claims.get("scope", "MISSING"),
}
@@ -205,7 +205,9 @@ class JoynChannel:
return json.dumps(self.to_dict(), indent=indent, ensure_ascii=False)
def __str__(self) -> str:
return f"JoynChannel(name='{self.name}', id='{self.channel_id}', type='{self.content_type}')"
return (
f"JoynChannel(name='{self.name}', id='{self.channel_id}', type='{self.content_type}')"
)
def __repr__(self) -> str:
return self.__str__()
@@ -15,20 +15,37 @@ from ...base.models.streaming_channel import StreamingChannel
from ...base.provider import AuthType, StreamingProvider
from ...base.utils.logger import logger
from .auth import JoynAuthenticator
from .constants import (CONTENT_TYPE_LIVE, CONTENT_TYPE_VOD,
COUNTRY_TENANT_MAPPING, DEFAULT_EPG_WINDOW_HOURS,
DEFAULT_LIVESTREAM_TYPES, DEFAULT_MAX_RETRIES,
DEFAULT_PLATFORM, DEFAULT_REQUEST_TIMEOUT,
DEFAULT_VIDEO_CONFIG, DRM_REQUEST_HEADERS,
DRM_SYSTEM_WIDEVINE, ERROR_CODES,
GRAPHQL_LIVE_CHANNELS_FILTER, GRAPHQL_MAX_RESULTS,
GRAPHQL_OFFSET, GRAPHQL_PERSISTED_QUERY_VERSION,
GRAPHQL_QUERY_HASHES, JOYN_API_BASE_HEADERS,
JOYN_CLIENT_VERSION, JOYN_DOMAINS,
JOYN_GRAPHQL_BASE_HEADERS, JOYN_GRAPHQL_ENDPOINTS,
JOYN_LOGO, JOYN_STREAMING_ENDPOINTS, JOYN_USER_AGENT,
MODE_LIVE, MODE_VOD, SIGNATURE_SECRET_KEY,
SUPPORTED_COUNTRIES)
from .constants import (
CONTENT_TYPE_LIVE,
CONTENT_TYPE_VOD,
COUNTRY_TENANT_MAPPING,
DEFAULT_EPG_WINDOW_HOURS,
DEFAULT_LIVESTREAM_TYPES,
DEFAULT_MAX_RETRIES,
DEFAULT_PLATFORM,
DEFAULT_REQUEST_TIMEOUT,
DEFAULT_VIDEO_CONFIG,
DRM_REQUEST_HEADERS,
DRM_SYSTEM_WIDEVINE,
ERROR_CODES,
GRAPHQL_LIVE_CHANNELS_FILTER,
GRAPHQL_MAX_RESULTS,
GRAPHQL_OFFSET,
GRAPHQL_PERSISTED_QUERY_VERSION,
GRAPHQL_QUERY_HASHES,
JOYN_API_BASE_HEADERS,
JOYN_CLIENT_VERSION,
JOYN_DOMAINS,
JOYN_GRAPHQL_BASE_HEADERS,
JOYN_GRAPHQL_ENDPOINTS,
JOYN_LOGO,
JOYN_STREAMING_ENDPOINTS,
JOYN_USER_AGENT,
MODE_LIVE,
MODE_VOD,
SIGNATURE_SECRET_KEY,
SUPPORTED_COUNTRIES,
)
from .models import JoynChannel, PlaybackRestrictedException
@@ -91,9 +108,7 @@ class JoynProvider(StreamingProvider):
"""
if not self.validate_country(country):
supported = ", ".join(self.SUPPORTED_COUNTRIES)
raise ValueError(
f"Unsupported country: {country}. " f"Joyn supports: {supported}"
)
raise ValueError(f"Unsupported country: {country}. " f"Joyn supports: {supported}")
super().__init__(country=country)
@@ -170,9 +185,7 @@ class JoynProvider(StreamingProvider):
)
return self.bearer_token
def get_dynamic_manifest_params(
self, channel: StreamingChannel, **kwargs
) -> Optional[str]:
def get_dynamic_manifest_params(self, channel: StreamingChannel, **kwargs) -> Optional[str]:
return None
def _get_graphql_headers(self) -> Dict[str, str]:
@@ -251,9 +264,7 @@ class JoynProvider(StreamingProvider):
if fetch_manifests and populate_streaming_data:
channels = self.populate_streaming_data(channels)
logger.info(
f"Successfully fetched {len(channels)} channels for country {self.country}"
)
logger.info(f"Successfully fetched {len(channels)} channels for country {self.country}")
return channels
except Exception as e:
@@ -278,9 +289,7 @@ class JoynProvider(StreamingProvider):
if "logo" in stream_data and "url" in stream_data["logo"]:
logo_url = stream_data["logo"]["url"]
content_type = (
CONTENT_TYPE_LIVE if stream_type == "LINEAR" else CONTENT_TYPE_VOD
)
content_type = CONTENT_TYPE_LIVE if stream_type == "LINEAR" else CONTENT_TYPE_VOD
mode = MODE_LIVE if stream_type == "LINEAR" else MODE_VOD
joyn_channel = JoynChannel(
@@ -314,9 +323,7 @@ class JoynProvider(StreamingProvider):
return channels
def get_entitlement_token(
self, content_id: str, content_type: str = CONTENT_TYPE_LIVE
) -> str:
def get_entitlement_token(self, content_id: str, content_type: str = CONTENT_TYPE_LIVE) -> str:
"""
Get entitlement token for content
@@ -364,9 +371,7 @@ class JoynProvider(StreamingProvider):
f"Playback restricted for {content_id}: {msg}"
)
else:
raise Exception(
f"Entitlement error for {content_id} ({code}): {msg}"
)
raise Exception(f"Entitlement error for {content_id} ({code}): {msg}")
except (json.JSONDecodeError, KeyError, IndexError) as e:
raise Exception(
f"Bad response for {content_id} (400), failed to parse error: {e}"
@@ -489,20 +494,14 @@ class JoynProvider(StreamingProvider):
except Exception as e:
retries += 1
if retries < max_retries:
logger.debug(
f"Retry {retries}/{max_retries} for {channel.name}: {e}"
)
logger.debug(f"Retry {retries}/{max_retries} for {channel.name}: {e}")
time.sleep(1)
else:
logger.error(
f"Failed to get streaming data for {channel.name}: {e}"
)
logger.error(f"Failed to get streaming data for {channel.name}: {e}")
logger.info(f"Streaming data population complete:")
logger.info(f" Successful: {len(successful_channels)}")
logger.info(
f" Restricted: {len([c for c in channels if c not in successful_channels])}"
)
logger.info(f" Restricted: {len([c for c in channels if c not in successful_channels])}")
logger.info(f" Total: {len(channels)}")
return successful_channels
@@ -589,9 +588,7 @@ class JoynProvider(StreamingProvider):
content_id=channel_id, content_type=content_type
)
playlist_data = self.get_channel_playlist(
channel_id, entitlement_token, video_config
)
playlist_data = self.get_channel_playlist(channel_id, entitlement_token, video_config)
return playlist_data.get("manifestUrl")
@@ -623,9 +620,7 @@ class JoynProvider(StreamingProvider):
content_id=channel_id, content_type=content_type
)
playlist_data = self.get_channel_playlist(
channel_id, entitlement_token, video_config
)
playlist_data = self.get_channel_playlist(channel_id, entitlement_token, video_config)
license_url = playlist_data.get("licenseUrl")
if not license_url:
@@ -683,9 +678,7 @@ class JoynProvider(StreamingProvider):
headers = self._get_graphql_headers()
# Joyn EPG would require additional GraphQL queries
logger.info(
f"EPG data requested for channel {channel_id} - not yet implemented"
)
logger.info(f"EPG data requested for channel {channel_id} - not yet implemented")
return []
except Exception as e:
@@ -1,13 +1,22 @@
# streaming_providers/providers/magenta2/__init__.py
from .auth import Magenta2Authenticator, Magenta2AuthToken, Magenta2Credentials
from .config_models import (BootstrapConfig, DrmConfig, ManifestConfig,
MpxConfig, OpenIDConfig, ProviderConfig,
TvHubConfig)
from .config_models import (
BootstrapConfig,
DrmConfig,
ManifestConfig,
MpxConfig,
OpenIDConfig,
ProviderConfig,
TvHubConfig,
)
from .constants import SUPPORTED_COUNTRIES
from .discovery import DiscoveryService
from .endpoint_manager import EndpointCategory, EndpointManager
from .models import (DeviceLimitExceededException, Magenta2Channel,
Magenta2PlaybackRestrictedException)
from .models import (
DeviceLimitExceededException,
Magenta2Channel,
Magenta2PlaybackRestrictedException,
)
from .provider import Magenta2Provider
__all__ = [
@@ -18,14 +18,21 @@ import uuid
from dataclasses import dataclass, field
from typing import Any, Dict, Optional
from ...base.auth.base_auth import (BaseAuthenticator, BaseAuthToken,
TokenAuthLevel)
from ...base.auth.base_auth import BaseAuthenticator, BaseAuthToken, TokenAuthLevel
from ...base.auth.credentials import ClientCredentials
from ...base.utils.logger import logger
from .constants import (APPVERSION2, DEFAULT_COUNTRY, DEFAULT_PLATFORM, IDM,
MAGENTA2_CLIENT_IDS, MAGENTA2_PLATFORMS,
SSO_USER_AGENT, SUPPORTED_COUNTRIES,
TAA_REQUEST_TEMPLATE)
from .constants import (
APPVERSION2,
DEFAULT_COUNTRY,
DEFAULT_PLATFORM,
IDM,
MAGENTA2_CLIENT_IDS,
MAGENTA2_PLATFORMS,
SSO_USER_AGENT,
SUPPORTED_COUNTRIES,
TAA_REQUEST_TEMPLATE,
)
# Import Magenta2-specific components
from .sam3_client import Sam3Client
from .sso_client import SsoClient
@@ -142,11 +149,7 @@ class Magenta2UserCredentials(Magenta2Credentials):
def validate_user_credentials(self) -> bool:
"""Validate user credentials"""
return (
self.has_user_credentials()
and len(self.username) > 0
and len(self.password) > 0
)
return self.has_user_credentials() and len(self.username) > 0 and len(self.password) > 0
@property
def credential_type(self) -> str:
@@ -258,9 +261,7 @@ class Magenta2AuthConfig:
return headers
@staticmethod
def get_sso_headers(
session_id: str = None, device_id: str = None
) -> Dict[str, str]:
def get_sso_headers(session_id: str = None, device_id: str = None) -> Dict[str, str]:
"""Get headers for SSO requests"""
headers = {
"User-Agent": SSO_USER_AGENT,
@@ -431,34 +432,22 @@ class Magenta2Authenticator(BaseAuthenticator):
# TokenFlowManager handles all token acquisition
return {}
def _create_token_from_response(
self, response_data: Dict[str, Any]
) -> BaseAuthToken:
def _create_token_from_response(self, response_data: Dict[str, Any]) -> BaseAuthToken:
"""
Create token object from API response and compose persona token
"""
# Handle different response key formats
access_token = response_data.get(
"access_token", response_data.get("accessToken")
)
access_token = response_data.get("access_token", response_data.get("accessToken"))
if not access_token:
raise ValueError("No access token in response")
# Create token with ALL fields
token = Magenta2AuthToken(
access_token=access_token,
refresh_token=response_data.get(
"refresh_token", response_data.get("refreshToken", "")
),
token_type=response_data.get(
"token_type", response_data.get("tokenType", "Bearer")
),
expires_in=response_data.get(
"expires_in", response_data.get("expiresIn", 3600)
),
issued_at=response_data.get(
"issued_at", response_data.get("issuedAt", time.time())
),
refresh_token=response_data.get("refresh_token", response_data.get("refreshToken", "")),
token_type=response_data.get("token_type", response_data.get("tokenType", "Bearer")),
expires_in=response_data.get("expires_in", response_data.get("expiresIn", 3600)),
issued_at=response_data.get("issued_at", response_data.get("issuedAt", time.time())),
# Magenta2-specific fields from JWT
dc_cts_persona_token=response_data.get("dc_cts_persona_token"),
persona_id=response_data.get("persona_id"),
@@ -495,9 +484,7 @@ class Magenta2Authenticator(BaseAuthenticator):
token.account_uri = f"urn:theplatform:auth:root:{self._mpx_account_pid}"
composed = token.compose_persona_token()
if composed:
logger.info(
"✓ Persona token composed using constructed account_uri"
)
logger.info("✓ Persona token composed using constructed account_uri")
# Classify token if it's NOT from line_auth
if not response_data.get("auth_source") == "line_auth":
@@ -558,22 +545,13 @@ class Magenta2Authenticator(BaseAuthenticator):
if not token or not token.access_token:
return TokenAuthLevel.UNKNOWN
claims = (
token.get_jwt_claims() if hasattr(token, "get_jwt_claims") else None
)
claims = token.get_jwt_claims() if hasattr(token, "get_jwt_claims") else None
if not claims:
# If we can't parse claims, check token attributes
if (
hasattr(token, "dc_cts_persona_token")
and token.dc_cts_persona_token
):
logger.debug(
"Token classified as USER_AUTHENTICATED (persona token present)"
)
if hasattr(token, "dc_cts_persona_token") and token.dc_cts_persona_token:
logger.debug("Token classified as USER_AUTHENTICATED (persona token present)")
return TokenAuthLevel.USER_AUTHENTICATED
logger.debug(
"Token classified as CLIENT_CREDENTIALS (no claims, no persona token)"
)
logger.debug("Token classified as CLIENT_CREDENTIALS (no claims, no persona token)")
return TokenAuthLevel.CLIENT_CREDENTIALS
logger.debug(f"JWT claims for classification: {list(claims.keys())}")
@@ -599,9 +577,7 @@ class Magenta2Authenticator(BaseAuthenticator):
for key in user_claim_keys:
if key in claims:
logger.debug(
f"Token classified as USER_AUTHENTICATED (found {key} in JWT)"
)
logger.debug(f"Token classified as USER_AUTHENTICATED (found {key} in JWT)")
return TokenAuthLevel.USER_AUTHENTICATED
# Check for client credentials patterns
@@ -631,9 +607,7 @@ class Magenta2Authenticator(BaseAuthenticator):
logger.info("Refreshing Magenta2 token via TokenFlowManager")
# Force refresh through TokenFlowManager
persona_result = self.token_flow_manager.get_persona_token(
force_refresh=True
)
persona_result = self.token_flow_manager.get_persona_token(force_refresh=True)
if persona_result.success:
token = Magenta2AuthToken(
@@ -678,16 +652,12 @@ class Magenta2Authenticator(BaseAuthenticator):
# Get QR code URL from dynamic endpoints
if "login_qr_code" in self._dynamic_endpoints:
qr_code_url_template = self._dynamic_endpoints["login_qr_code"]
logger.debug(
f"QR code URL from dynamic endpoints: {qr_code_url_template}"
)
logger.debug(f"QR code URL from dynamic endpoints: {qr_code_url_template}")
if self._openid_config:
issuer_url = self._openid_config.get("issuer")
oauth_endpoint = self._openid_config.get("token_endpoint")
backchannel_start_url = self._openid_config.get(
"backchannel_auth_start"
)
backchannel_start_url = self._openid_config.get("backchannel_auth_start")
self._sam3_client = Sam3Client(
http_manager=self._http_manager,
@@ -715,9 +685,7 @@ class Magenta2Authenticator(BaseAuthenticator):
def _initialize_taa_client(self) -> None:
"""Initialize TAA client"""
self._taa_client = TaaClient(
http_manager=self._http_manager, platform=self.platform
)
self._taa_client = TaaClient(http_manager=self._http_manager, platform=self.platform)
logger.debug("TAA client initialized")
def _initialize_token_flow_manager(self) -> None:
@@ -725,9 +693,7 @@ class Magenta2Authenticator(BaseAuthenticator):
if self._sam3_client and self._taa_client:
session_manager = getattr(self.settings_manager, "session_manager", None)
if not session_manager:
logger.error(
"Cannot initialize TokenFlowManager: No session_manager available"
)
logger.error("Cannot initialize TokenFlowManager: No session_manager available")
return
self.token_flow_manager = TokenFlowManager(
@@ -765,9 +731,7 @@ class Magenta2Authenticator(BaseAuthenticator):
if not self.token_flow_manager:
raise Exception("TokenFlowManager not initialized")
persona_result = self.token_flow_manager.get_persona_token(
force_refresh=force_refresh
)
persona_result = self.token_flow_manager.get_persona_token(force_refresh=force_refresh)
if not persona_result.success:
raise Exception(f"Failed to get persona token: {persona_result.error}")
@@ -808,13 +772,9 @@ class Magenta2Authenticator(BaseAuthenticator):
if self._sam3_client:
self._sam3_client.update_sam3_client_id(client_id)
logger.info(
f"✓ Updated SAM3 client ID: {old_client_id[:8]}... -> {client_id[:8]}..."
)
logger.info(f"✓ Updated SAM3 client ID: {old_client_id[:8]}... -> {client_id[:8]}...")
else:
logger.debug(
f"Updated SAM3 client ID (no client to update yet): {client_id}"
)
logger.debug(f"Updated SAM3 client ID (no client to update yet): {client_id}")
def update_client_model(self, client_model: str) -> None:
"""Update client model"""
@@ -846,9 +806,7 @@ class Magenta2Authenticator(BaseAuthenticator):
self._mpx_account_pid = account_pid
logger.debug(f"MPX account PID set: {account_pid}")
def set_device_token(
self, device_token: str, authorize_tokens_url: str = None
) -> None:
def set_device_token(self, device_token: str, authorize_tokens_url: str = None) -> None:
"""
Enhanced device token setup with both endpoints
@@ -863,9 +821,7 @@ class Magenta2Authenticator(BaseAuthenticator):
if self._sam3_client and authorize_tokens_url:
self._sam3_client.line_auth_endpoint = authorize_tokens_url
self._sam3_client.token_endpoint = authorize_tokens_url # Backwards compat
logger.info(
f"✓ Updated SAM3 client with line auth endpoint: {authorize_tokens_url}"
)
logger.info(f"✓ Updated SAM3 client with line auth endpoint: {authorize_tokens_url}")
logger.debug("Device token configured with line authentication support")
@@ -931,9 +887,7 @@ class Magenta2Authenticator(BaseAuthenticator):
def can_use_remote_login(self) -> bool:
"""Check if remote login components are available"""
return (
self._sam3_client is not None and self._sam3_client.can_use_remote_login()
)
return self._sam3_client is not None and self._sam3_client.can_use_remote_login()
def get_authentication_capabilities(self) -> Dict[str, Any]:
"""Get authentication capabilities information"""
@@ -943,9 +897,7 @@ class Magenta2Authenticator(BaseAuthenticator):
return {
"line_auth_available": line_auth_available,
"remote_login_available": remote_login_available,
"user_credentials_available": isinstance(
self.credentials, Magenta2UserCredentials
)
"user_credentials_available": isinstance(self.credentials, Magenta2UserCredentials)
and self.credentials.has_user_credentials(),
"client_credentials_available": True,
"preferred_flow": (
@@ -969,9 +921,7 @@ class Magenta2Authenticator(BaseAuthenticator):
def get_authentication_flow_info(self) -> Dict[str, Any]:
"""Get authentication flow information"""
base_info = {
"user_credentials_available": isinstance(
self.credentials, Magenta2UserCredentials
)
"user_credentials_available": isinstance(self.credentials, Magenta2UserCredentials)
and self.credentials.has_user_credentials(),
"client_credentials_available": True,
"sam3_client_available": self._sam3_client is not None,
@@ -991,9 +941,7 @@ class Magenta2Authenticator(BaseAuthenticator):
# Add TAA-specific info if available
if self._taa_client and self._current_token:
base_info["taa_token_valid"] = self.validate_taa_token(
self._current_token.access_token
)
base_info["taa_token_valid"] = self.validate_taa_token(self._current_token.access_token)
return base_info
@@ -1021,9 +969,7 @@ class Magenta2Authenticator(BaseAuthenticator):
"""
try:
if not self._device_token or not self._authorize_tokens_url:
logger.warning(
"Line auth skipped - missing device token or authorize URL"
)
logger.warning("Line auth skipped - missing device token or authorize URL")
return False
if not self._sam3_client:
@@ -1070,9 +1016,7 @@ class Magenta2Authenticator(BaseAuthenticator):
if response_data and "refresh_token" in response_data:
self._line_auth_refresh_token = response_data["refresh_token"]
logger.info(f"✓ Stored refresh token from line auth for SAM3 requests")
logger.debug(
f"Refresh token preview: {self._line_auth_refresh_token[:20]}..."
)
logger.debug(f"Refresh token preview: {self._line_auth_refresh_token[:20]}...")
return response_data
@@ -1150,13 +1094,9 @@ class Magenta2Authenticator(BaseAuthenticator):
return {
"has_access_token": bool(token.access_token),
"has_dc_cts_persona_token": bool(
getattr(token, "dc_cts_persona_token", None)
),
"has_dc_cts_persona_token": bool(getattr(token, "dc_cts_persona_token", None)),
"has_account_uri": bool(getattr(token, "account_uri", None)),
"has_composed_persona_token": bool(
getattr(token, "composed_persona_token", None)
),
"has_composed_persona_token": bool(getattr(token, "composed_persona_token", None)),
"persona_token_preview": (
getattr(token, "composed_persona_token", "")[:50] + "..."
if getattr(token, "composed_persona_token", None)
@@ -1189,15 +1129,11 @@ class Magenta2Authenticator(BaseAuthenticator):
"auth_level": self._current_token.auth_level.value,
"is_expired": self._current_token.is_expired,
"has_refresh": bool(self._current_token.refresh_token),
"has_persona_token": bool(
getattr(self._current_token, "dc_cts_persona_token", None)
),
"has_persona_token": bool(getattr(self._current_token, "dc_cts_persona_token", None)),
"jwt_claims_available": bool(claims),
"key_claims": (
{
"client_id": claims.get(
"client_id", claims.get("clientId", "MISSING")
),
"client_id": claims.get("client_id", claims.get("clientId", "MISSING")),
"persona_id": claims.get(
"dc_cts_personaId", claims.get("personaId", "MISSING")
),
@@ -50,9 +50,7 @@ def extract_and_release_lock(
return False
# Build release URL with the same client_id formatted as player_{client_id}
base_url = (
lock_params["concurrencyServiceUrl"].rstrip("/") + "/web/Concurrency/unlock"
)
base_url = lock_params["concurrencyServiceUrl"].rstrip("/") + "/web/Concurrency/unlock"
formatted_client_id = f"player_{client_id}"
params = {
@@ -107,9 +105,7 @@ def extract_and_release_lock(
)
return False
else:
logger.warning(
f"Concurrency lock release failed with status: {response.status_code}"
)
logger.warning(f"Concurrency lock release failed with status: {response.status_code}")
return False
except Exception as e:
@@ -22,9 +22,7 @@ class BootstrapConfig:
raw_data: Dict[str, Any] = field(default_factory=dict)
@classmethod
def from_api_response(
cls, bootstrap_data: Dict[str, Any], platform: str
) -> "BootstrapConfig":
def from_api_response(cls, bootstrap_data: Dict[str, Any], platform: str) -> "BootstrapConfig":
"""Create BootstrapConfig from API response"""
base_settings = bootstrap_data.get("baseSettings", {})
@@ -123,27 +121,20 @@ class MpxConfig:
for feed_key in feed_keys:
feed_url = get_param(feed_key)
if feed_url:
simple_key = (
feed_key.replace("mpx", "")
.replace("Url", "")
.replace("BasicUrl", "")
)
simple_key = feed_key.replace("mpx", "").replace("Url", "").replace("BasicUrl", "")
feeds[simple_key] = feed_url
# ADD THIS: Extract channel stations feed
channel_stations_feed = get_param("mpxDefaultUrlAllChannelStationsFeed")
return cls(
account_pid=get_param("mpxAccountPid")
or mpx_data.get("accountPid", "mdeprod"),
account_pid=get_param("mpxAccountPid") or mpx_data.get("accountPid", "mdeprod"),
license_service_url=get_param("mpxBasicUrlGetApplicableDistributionRights")
or mpx_data.get("licenseServiceUrl", ""),
selector_service_url=get_param("mpxBasicUrlSelectorService")
or mpx_data.get("selectorServiceUrl", ""),
user_profile_url=get_param("mpxUserProfileUrl")
or mpx_data.get("userProfileUrl"),
bookmark_base_url=get_param("mpxBookmarkBaseUrl")
or mpx_data.get("bookmarkBaseUrl"),
user_profile_url=get_param("mpxUserProfileUrl") or mpx_data.get("userProfileUrl"),
bookmark_base_url=get_param("mpxBookmarkBaseUrl") or mpx_data.get("bookmarkBaseUrl"),
pvr_base_url=get_param("mpxPvrBaseUrl") or mpx_data.get("pvrBaseUrl"),
feeds=feeds,
channel_stations_feed=channel_stations_feed,
@@ -239,9 +230,7 @@ class DrmConfig:
else:
vod_widevine = None # Not in parameters array
logger.debug(
f"DRM config: widevine={bool(widevine_url)}, fairplay={bool(fairplay_url)}"
)
logger.debug(f"DRM config: widevine={bool(widevine_url)}, fairplay={bool(fairplay_url)}")
return cls(
widevine_license_url=widevine_url or "",
@@ -130,9 +130,7 @@ MAGENTA2_FALLBACK_ENDPOINTS = {
"ENTITLEMENT": "https://entitlement.p7s1.io/api/user/entitlement-token",
}
MAGENTA2_FALLBACK_ACCOUNT_URI = (
"http://access.auth.theplatform.com/data/Account/2709353023"
)
MAGENTA2_FALLBACK_ACCOUNT_URI = "http://access.auth.theplatform.com/data/Account/2709353023"
# ============================================================================
# Application Configuration
@@ -5,11 +5,15 @@ from typing import Any, Dict, Optional
from ...base.models.proxy_models import ProxyConfig
from ...base.network import HTTPManager
from ...base.utils.logger import logger
from .config_models import (BootstrapConfig, ManifestConfig, OpenIDConfig,
ProviderConfig)
from .constants import (BOOTSTRAP_CACHE_DURATION, DEFAULT_REQUEST_TIMEOUT,
MAGENTA2_BOOTSTRAP_URL, MAGENTA2_MANIFEST_URL,
OPENID_CONFIG_CACHE_DURATION, SUBSCRIBER_TYPES)
from .config_models import BootstrapConfig, ManifestConfig, OpenIDConfig, ProviderConfig
from .constants import (
BOOTSTRAP_CACHE_DURATION,
DEFAULT_REQUEST_TIMEOUT,
MAGENTA2_BOOTSTRAP_URL,
MAGENTA2_MANIFEST_URL,
OPENID_CONFIG_CACHE_DURATION,
SUBSCRIBER_TYPES,
)
class DiscoveryService:
@@ -97,9 +101,7 @@ class DiscoveryService:
self._create_fallback_configuration()
raise
def discover_bootstrap(
self, force_refresh: bool = False
) -> Optional[BootstrapConfig]:
def discover_bootstrap(self, force_refresh: bool = False) -> Optional[BootstrapConfig]:
"""
Discover bootstrap configuration
@@ -152,9 +154,7 @@ class DiscoveryService:
if self._bootstrap_config.taa_url:
logger.debug(f"TAA URL: {self._bootstrap_config.taa_url}")
if self._bootstrap_config.device_tokens_url:
logger.debug(
f"Device tokens URL: {self._bootstrap_config.device_tokens_url}"
)
logger.debug(f"Device tokens URL: {self._bootstrap_config.device_tokens_url}")
return self._bootstrap_config
@@ -165,9 +165,7 @@ class DiscoveryService:
self._last_bootstrap = None
return None
def discover_manifest(
self, force_refresh: bool = False
) -> Optional[ManifestConfig]:
def discover_manifest(self, force_refresh: bool = False) -> Optional[ManifestConfig]:
"""
ENHANCED: Discover manifest configuration including device token
Uses correct manifest endpoint parameters
@@ -188,28 +186,27 @@ class DiscoveryService:
# Determine manifest URL - prefer device_tokens_url from bootstrap
if self._bootstrap_config and self._bootstrap_config.device_tokens_url:
manifest_url = self._bootstrap_config.device_tokens_url
logger.debug(
f"Using bootstrap device_tokens_url for manifest: {manifest_url}"
)
logger.debug(f"Using bootstrap device_tokens_url for manifest: {manifest_url}")
else:
terminal_type = self.terminal_type.lower().replace("_", "-")
manifest_url = MAGENTA2_MANIFEST_URL.format(terminal_type=terminal_type)
logger.debug(f"Using fallback manifest URL: {manifest_url}")
# Build correct manifest parameters
from .constants import (MAGENTA2_APP_NAME, MAGENTA2_APP_VERSION,
MAGENTA2_RUNTIME_VERSION,
MANIFEST_FIRMWARE_MAPPINGS,
MANIFEST_MODEL_MAPPINGS)
from .constants import (
MAGENTA2_APP_NAME,
MAGENTA2_APP_VERSION,
MAGENTA2_RUNTIME_VERSION,
MANIFEST_FIRMWARE_MAPPINGS,
MANIFEST_MODEL_MAPPINGS,
)
params = {
"model": MANIFEST_MODEL_MAPPINGS.get(self.platform, "DT:ATV-AndroidTV"),
"deviceId": self.device_id,
"appname": MAGENTA2_APP_NAME,
"appVersion": MAGENTA2_APP_VERSION,
"firmware": MANIFEST_FIRMWARE_MAPPINGS.get(
self.platform, "API level 30"
),
"firmware": MANIFEST_FIRMWARE_MAPPINGS.get(self.platform, "API level 30"),
"runtimeVersion": MAGENTA2_RUNTIME_VERSION,
"duid": self.device_id, # Same as deviceId
}
@@ -237,9 +234,7 @@ class DiscoveryService:
if device_token:
logger.info("✓ Device token found in manifest")
logger.debug(
f"Device token preview: {device_token[:20]}...{device_token[-10:]}"
)
logger.debug(f"Device token preview: {device_token[:20]}...{device_token[-10:]}")
else:
logger.warning("⚠️ Device token NOT found in manifest response")
@@ -268,9 +263,7 @@ class DiscoveryService:
self._last_manifest = None
return None
def discover_openid_config(
self, force_refresh: bool = False
) -> Optional[OpenIDConfig]:
def discover_openid_config(self, force_refresh: bool = False) -> Optional[OpenIDConfig]:
"""
Discover OpenID Connect configuration
@@ -285,10 +278,7 @@ class DiscoveryService:
return self._openid_config
try:
if (
not self._bootstrap_config
or not self._bootstrap_config.openid_config_url
):
if not self._bootstrap_config or not self._bootstrap_config.openid_config_url:
logger.warning("No OpenID config URL available from bootstrap")
return None
@@ -367,14 +357,10 @@ class DiscoveryService:
"mpx_account_uri": self._manifest_config.mpx.get_account_uri(),
"feed_count": len(self._manifest_config.mpx.feeds),
"has_device_token": bool(device_token),
"device_token_preview": (
device_token[:20] + "..." if device_token else None
),
"device_token_preview": (device_token[:20] + "..." if device_token else None),
"drm_endpoints": {
"widevine": bool(self._manifest_config.drm.widevine_license_url),
"vod_widevine": bool(
self._manifest_config.drm.vod_widevine_license_url
),
"vod_widevine": bool(self._manifest_config.drm.vod_widevine_license_url),
"fairplay": bool(self._manifest_config.drm.fairplay_license_url),
},
"tvhub_count": len(self._manifest_config.tv_hubs.base_urls),
@@ -383,17 +369,13 @@ class DiscoveryService:
if self._openid_config:
status["openid"] = {
"has_token_endpoint": bool(self._openid_config.token_endpoint),
"has_authorization_endpoint": bool(
self._openid_config.authorization_endpoint
),
"has_authorization_endpoint": bool(self._openid_config.authorization_endpoint),
}
# Cache status
now = time.time()
status["cache"] = {
"bootstrap_age": (
now - self._last_bootstrap if self._last_bootstrap else None
),
"bootstrap_age": (now - self._last_bootstrap if self._last_bootstrap else None),
"manifest_age": now - self._last_manifest if self._last_manifest else None,
"openid_age": now - self._last_openid if self._last_openid else None,
}
@@ -57,9 +57,7 @@ class EndpointManager:
# Authentication endpoints
if bootstrap.taa_url:
self._add_endpoint(
"taa_auth", EndpointCategory.AUTHENTICATION, bootstrap.taa_url
)
self._add_endpoint("taa_auth", EndpointCategory.AUTHENTICATION, bootstrap.taa_url)
if bootstrap.openid_config_url:
self._add_endpoint(
@@ -94,9 +92,7 @@ class EndpointManager:
)
if bootstrap.account_base_url:
self._add_endpoint(
"account_base", EndpointCategory.USER, bootstrap.account_base_url
)
self._add_endpoint("account_base", EndpointCategory.USER, bootstrap.account_base_url)
if bootstrap.consumer_accounts_url:
self._add_endpoint(
@@ -129,9 +125,7 @@ class EndpointManager:
EndpointCategory.CONTENT,
manifest.mpx.channel_stations_feed,
)
logger.info(
f"Channel stations feed found: {manifest.mpx.channel_stations_feed}"
)
logger.info(f"Channel stations feed found: {manifest.mpx.channel_stations_feed}")
# DRM endpoints
if manifest.drm.widevine_license_url:
@@ -159,17 +153,13 @@ class EndpointManager:
for feed_name, feed_template in manifest.mpx.feeds.items():
resolved_url = self.config.get_resolved_feed_url(feed_name)
if resolved_url:
self._add_endpoint(
f"mpx_feed_{feed_name}", EndpointCategory.MPX, resolved_url
)
self._add_endpoint(f"mpx_feed_{feed_name}", EndpointCategory.MPX, resolved_url)
# TV Hub URLs (resolved with client model)
for hub_name in manifest.tv_hubs.base_urls.keys():
resolved_url = self.config.get_resolved_tvhub_url(hub_name)
if resolved_url:
self._add_endpoint(
f"tvhub_{hub_name}", EndpointCategory.TVHUBS, resolved_url
)
self._add_endpoint(f"tvhub_{hub_name}", EndpointCategory.TVHUBS, resolved_url)
def _add_openid_endpoints(self) -> None:
"""Add endpoints from OpenID configuration"""
@@ -260,9 +250,7 @@ class EndpointManager:
def get_endpoints_by_category(self, category: EndpointCategory) -> Dict[str, str]:
"""Get all endpoints for a specific category"""
return {
name: info.url
for name, info in self._endpoints.items()
if info.category == category
name: info.url for name, info in self._endpoints.items() if info.category == category
}
def has_endpoint(self, name: str) -> bool:
@@ -15,15 +15,26 @@ from ...base.models.streaming_channel import StreamingChannel
from ...base.network import HTTPManagerFactory, ProxyConfigManager
from ...base.provider import StreamingProvider
from ...base.utils.logger import logger
from .auth import (Magenta2Authenticator, Magenta2Credentials,
Magenta2UserCredentials)
from .auth import Magenta2Authenticator, Magenta2Credentials, Magenta2UserCredentials
from .config_models import ProviderConfig
from .constants import (CONTENT_TYPE_LIVE, CONTENT_TYPE_VOD, DEFAULT_COUNTRY,
DEFAULT_EPG_WINDOW_HOURS, DEFAULT_MAX_RETRIES,
DEFAULT_PLATFORM, DEFAULT_REQUEST_TIMEOUT,
DRM_REQUEST_HEADERS, DRM_SYSTEM_WIDEVINE, ERROR_CODES,
MAGENTA2_CLIENT_IDS, MAGENTA2_LOGO, MAGENTA2_PLATFORMS,
MODE_LIVE, MODE_VOD, SUPPORTED_COUNTRIES)
from .constants import (
CONTENT_TYPE_LIVE,
CONTENT_TYPE_VOD,
DEFAULT_COUNTRY,
DEFAULT_EPG_WINDOW_HOURS,
DEFAULT_MAX_RETRIES,
DEFAULT_PLATFORM,
DEFAULT_REQUEST_TIMEOUT,
DRM_REQUEST_HEADERS,
DRM_SYSTEM_WIDEVINE,
ERROR_CODES,
MAGENTA2_CLIENT_IDS,
MAGENTA2_LOGO,
MAGENTA2_PLATFORMS,
MODE_LIVE,
MODE_VOD,
SUPPORTED_COUNTRIES,
)
from .discovery import DiscoveryService
from .endpoint_manager import EndpointManager
from .models import Magenta2Channel, Magenta2PlaybackRestrictedException
@@ -155,9 +166,7 @@ class Magenta2Provider(StreamingProvider):
# 🚨 UPDATE AUTHENTICATOR WITH DISCOVERED CONFIG
if self.provider_config:
# Update authenticator with discovered client_id and models
self.authenticator.provider_config = (
self.provider_config
) # ✅ Store the config
self.authenticator.provider_config = self.provider_config # ✅ Store the config
logger.info("✓ ProviderConfig stored in authenticator")
# Also update TokenFlowManager if it exists
@@ -165,9 +174,7 @@ class Magenta2Provider(StreamingProvider):
hasattr(self.authenticator, "token_flow_manager")
and self.authenticator.token_flow_manager
):
self.authenticator.token_flow_manager.provider_config = (
self.provider_config
)
self.authenticator.token_flow_manager.provider_config = self.provider_config
logger.info("✓ ProviderConfig also stored in TokenFlowManager")
if self.provider_config.bootstrap.sam3_client_id:
@@ -195,9 +202,7 @@ class Magenta2Provider(StreamingProvider):
self.provider_config.bootstrap.client_model
)
else:
self.authenticator._client_model = (
self.provider_config.bootstrap.client_model
)
self.authenticator._client_model = self.provider_config.bootstrap.client_model
logger.debug(
f"Updated authenticator with client model: {self.provider_config.bootstrap.client_model}"
)
@@ -208,9 +213,7 @@ class Magenta2Provider(StreamingProvider):
self.provider_config.bootstrap.device_model
)
else:
self.authenticator._device_model = (
self.provider_config.bootstrap.device_model
)
self.authenticator._device_model = self.provider_config.bootstrap.device_model
logger.debug(
f"Updated authenticator with device model: {self.provider_config.bootstrap.device_model}"
)
@@ -221,9 +224,7 @@ class Magenta2Provider(StreamingProvider):
authorize_tokens_url = self.provider_config.get_authorize_tokens_url()
if device_token:
self.authenticator.set_device_token(
device_token, authorize_tokens_url
)
self.authenticator.set_device_token(device_token, authorize_tokens_url)
logger.debug("Device token configured in authenticator")
# CRITICAL: Pass MPX account PID for account URI construction
@@ -237,9 +238,7 @@ class Magenta2Provider(StreamingProvider):
# Pass OpenID configuration if available
if self.provider_config.openid:
self.authenticator.set_openid_config(
self.provider_config.openid.raw_data
)
self.authenticator.set_openid_config(self.provider_config.openid.raw_data)
# Update authenticator with discovered endpoints using public methods
if self.endpoint_manager:
@@ -252,14 +251,10 @@ class Magenta2Provider(StreamingProvider):
# Update authenticator with discovered endpoints using public method
if hasattr(self.authenticator, "update_dynamic_endpoints"):
self.authenticator.update_dynamic_endpoints(all_endpoints)
logger.info(
f"✓ Updated authenticator with {len(all_endpoints)} endpoints"
)
logger.info(f"✓ Updated authenticator with {len(all_endpoints)} endpoints")
elif hasattr(self.authenticator, "update_endpoints"):
self.authenticator.update_endpoints(all_endpoints)
logger.info(
f"✓ Updated authenticator with {len(all_endpoints)} endpoints"
)
logger.info(f"✓ Updated authenticator with {len(all_endpoints)} endpoints")
else:
logger.warning("No public method available to update endpoints")
@@ -268,20 +263,14 @@ class Magenta2Provider(StreamingProvider):
if qr_url and hasattr(self.authenticator, "update_sam3_qr_code_url"):
success = self.authenticator.update_sam3_qr_code_url(qr_url)
if success:
logger.info(
"✓ Successfully updated SAM3 client with QR code URL"
)
logger.info("✓ Successfully updated SAM3 client with QR code URL")
else:
logger.warning(
"✗ Failed to update SAM3 client with QR code URL"
)
logger.warning("✗ Failed to update SAM3 client with QR code URL")
# Initialize auth tokens (lazy - populated on first use)
self.device_token = None
self._persona_cache: Optional[PersonaResult] = None
self._smil_cache: Dict[str, Tuple[float, Dict]] = (
{}
) # channel_id -> (timestamp, smil_data)
self._smil_cache: Dict[str, Tuple[float, Dict]] = {} # channel_id -> (timestamp, smil_data)
self._smil_cache_ttl = 3600
logger.info("Magenta2 provider initialization completed successfully")
@@ -303,9 +292,7 @@ class Magenta2Provider(StreamingProvider):
"""Generate call ID for requests"""
return str(uuid.uuid4())
def _load_proxy_from_manager(
self, config_dir: Optional[str]
) -> Optional[ProxyConfig]:
def _load_proxy_from_manager(self, config_dir: Optional[str]) -> Optional[ProxyConfig]:
"""Load proxy configuration from ProxyConfigManager"""
try:
proxy_manager = ProxyConfigManager(config_dir)
@@ -325,9 +312,7 @@ class Magenta2Provider(StreamingProvider):
self.provider_config = self.discovery_service.discover_provider_config()
if not self.provider_config or not self.provider_config.is_complete:
logger.warning(
"Configuration discovery incomplete, some features may not work"
)
logger.warning("Configuration discovery incomplete, some features may not work")
# Initialize endpoint manager with discovered configuration
self.endpoint_manager = EndpointManager(self.provider_config)
@@ -336,21 +321,15 @@ class Magenta2Provider(StreamingProvider):
if self.endpoint_manager:
qr_url = self.endpoint_manager.get_endpoint("login_qr_code")
if qr_url:
logger.info(
f"✓ QR code endpoint discovered in endpoint manager: {qr_url}"
)
logger.info(f"✓ QR code endpoint discovered in endpoint manager: {qr_url}")
# PROPER FIX: Use public method to update SAM3 client
if hasattr(self.authenticator, "update_sam3_qr_code_url"):
success = self.authenticator.update_sam3_qr_code_url(qr_url)
if success:
logger.info(
"✓ Successfully updated SAM3 client with QR code URL"
)
logger.info("✓ Successfully updated SAM3 client with QR code URL")
else:
logger.warning(
"✗ Failed to update SAM3 client with QR code URL"
)
logger.warning("✗ Failed to update SAM3 client with QR code URL")
# Also debug the current status
if hasattr(self.authenticator, "get_sam3_client_status"):
@@ -364,16 +343,12 @@ class Magenta2Provider(StreamingProvider):
authorize_tokens_url = self.provider_config.get_authorize_tokens_url()
if device_token:
logger.info(
f"✓ Device token discovered (length: {len(device_token)})"
)
logger.info(f"✓ Device token discovered (length: {len(device_token)})")
else:
logger.warning("⚠️ No device token found in manifest")
if authorize_tokens_url:
logger.info(
f"✓ Line auth endpoint discovered: {authorize_tokens_url}"
)
logger.info(f"✓ Line auth endpoint discovered: {authorize_tokens_url}")
else:
logger.warning("⚠️ No authorize tokens URL found in manifest")
@@ -504,9 +479,7 @@ class Magenta2Provider(StreamingProvider):
try:
logger.info("Refreshing provider configuration")
new_config = self.discovery_service.discover_provider_config(
force_refresh=force
)
new_config = self.discovery_service.discover_provider_config(force_refresh=force)
if new_config and new_config.is_complete:
self.provider_config = new_config
@@ -515,18 +488,12 @@ class Magenta2Provider(StreamingProvider):
# Update authenticator with new config
if new_config.manifest:
device_token = new_config.manifest.raw_data.get("deviceToken")
authorize_tokens_url = new_config.manifest.raw_data.get(
"authorizeTokensUrl"
)
authorize_tokens_url = new_config.manifest.raw_data.get("authorizeTokensUrl")
if device_token:
self.authenticator.set_device_token(
device_token, authorize_tokens_url
)
self.authenticator.set_device_token(device_token, authorize_tokens_url)
if new_config.manifest.mpx.account_pid:
self.authenticator.set_mpx_account_pid(
new_config.manifest.mpx.account_pid
)
self.authenticator.set_mpx_account_pid(new_config.manifest.mpx.account_pid)
logger.info("Configuration refresh successful")
return True
@@ -560,9 +527,7 @@ class Magenta2Provider(StreamingProvider):
logger.warning("Device registration failed")
return False
else:
logger.warning(
"Device authentication not supported in current authenticator"
)
logger.warning("Device authentication not supported in current authenticator")
return False
except Exception as e:
@@ -610,9 +575,7 @@ class Magenta2Provider(StreamingProvider):
self._persona_cache = None
logger.debug("Cleared in-memory persona cache")
def get_dynamic_manifest_params(
self, channel: StreamingChannel, **kwargs
) -> Optional[str]:
def get_dynamic_manifest_params(self, channel: StreamingChannel, **kwargs) -> Optional[str]:
return None
@staticmethod
@@ -682,9 +645,7 @@ class Magenta2Provider(StreamingProvider):
}
# Build query string
query_string = "&".join(
[f"{k}={self._url_encode(v)}" for k, v in params.items()]
)
query_string = "&".join([f"{k}={self._url_encode(v)}" for k, v in params.items()])
return f"{base_url}/iss?{query_string}"
@@ -711,9 +672,7 @@ class Magenta2Provider(StreamingProvider):
if display_number is None:
# Process immediately if no number (no filtering needed)
channel = self._create_channel_from_entry(
entry, station_info, display_number
)
channel = self._create_channel_from_entry(entry, station_info, display_number)
if channel:
channels.append(channel)
continue
@@ -742,9 +701,7 @@ class Magenta2Provider(StreamingProvider):
# Convert best entries to channels
for display_number, (entry, station_info, _) in best_entries.items():
channel = self._create_channel_from_entry(
entry, station_info, display_number
)
channel = self._create_channel_from_entry(entry, station_info, display_number)
if channel:
channels.append(channel)
@@ -812,12 +769,8 @@ class Magenta2Provider(StreamingProvider):
url = self.endpoint_manager.get_endpoint("channel_stations")
if not url:
url = self.endpoint_manager.get_endpoint("channel_list")
if not url and self.endpoint_manager.has_endpoint(
"mpx_feed_entitledChannelsFeed"
):
url = self.endpoint_manager.get_endpoint(
"mpx_feed_entitledChannelsFeed"
)
if not url and self.endpoint_manager.has_endpoint("mpx_feed_entitledChannelsFeed"):
url = self.endpoint_manager.get_endpoint("mpx_feed_entitledChannelsFeed")
# Final fallback
if not url:
@@ -884,9 +837,7 @@ class Magenta2Provider(StreamingProvider):
for entry in entries:
try:
title = entry.get("title", "Unknown")
channel_id = (
entry.get("guid", "").split("/")[-1] if entry.get("guid") else ""
)
channel_id = entry.get("guid", "").split("/")[-1] if entry.get("guid") else ""
if not channel_id:
continue
@@ -899,9 +850,7 @@ class Magenta2Provider(StreamingProvider):
raw_data=entry,
)
streaming_channel = channel.to_streaming_channel(
provider_name=self.provider_name
)
streaming_channel = channel.to_streaming_channel(provider_name=self.provider_name)
channels.append(streaming_channel)
except Exception as e:
@@ -933,18 +882,13 @@ class Magenta2Provider(StreamingProvider):
logger.warning("Channel missing ID, skipping")
continue
title = channel_data.get(
"title", channel_data.get("name", "Unknown Channel")
)
title = channel_data.get("title", channel_data.get("name", "Unknown Channel"))
stream_type = channel_data.get("type", "LIVE")
quality = channel_data.get("quality", "")
logo_url = None
if "logo" in channel_data:
if (
isinstance(channel_data["logo"], dict)
and "url" in channel_data["logo"]
):
if isinstance(channel_data["logo"], dict) and "url" in channel_data["logo"]:
logo_url = channel_data["logo"]["url"]
elif isinstance(channel_data["logo"], str):
logo_url = channel_data["logo"]
@@ -952,9 +896,7 @@ class Magenta2Provider(StreamingProvider):
logo_url = channel_data["image"]
content_type = (
CONTENT_TYPE_LIVE
if stream_type.upper() == "LIVE"
else CONTENT_TYPE_VOD
CONTENT_TYPE_LIVE if stream_type.upper() == "LIVE" else CONTENT_TYPE_VOD
)
mode = MODE_LIVE if stream_type.upper() == "LIVE" else MODE_VOD
@@ -981,9 +923,7 @@ class Magenta2Provider(StreamingProvider):
return channels
def get_entitlement_token(
self, content_id: str, content_type: str = CONTENT_TYPE_LIVE
) -> str:
def get_entitlement_token(self, content_id: str, content_type: str = CONTENT_TYPE_LIVE) -> str:
"""
Get entitlement token using persona_token Basic auth
"""
@@ -1002,9 +942,7 @@ class Magenta2Provider(StreamingProvider):
)
try:
logger.debug(
f"Requesting entitlement token with persona token for: {content_id}"
)
logger.debug(f"Requesting entitlement token with persona token for: {content_id}")
response = self.http_manager.post(
url,
operation="auth",
@@ -1016,25 +954,19 @@ class Magenta2Provider(StreamingProvider):
if response.status_code == 400:
try:
error_data = response.json()
error_list = (
error_data if isinstance(error_data, list) else [error_data]
)
error_list = error_data if isinstance(error_data, list) else [error_data]
if len(error_list) > 0:
error = error_list[0]
code = error.get("code", error.get("errorCode", "UNKNOWN"))
msg = error.get(
"msg", error.get("message", "No error message provided")
)
msg = error.get("msg", error.get("message", "No error message provided"))
if code == ERROR_CODES["PLAYBACK_RESTRICTED"]:
raise Magenta2PlaybackRestrictedException(
f"Playback restricted for {content_id}: {msg}"
)
else:
raise Exception(
f"Entitlement error for {content_id} ({code}): {msg}"
)
raise Exception(f"Entitlement error for {content_id} ({code}): {msg}")
except (json.JSONDecodeError, KeyError, IndexError) as e:
raise Exception(
f"Bad response for {content_id} (400), failed to parse error: {e}"
@@ -1056,22 +988,16 @@ class Magenta2Provider(StreamingProvider):
raise
except KeyError as e:
logger.error(f"No entitlement token in response for {content_id}: {e}")
logger.debug(
f"Auth state: {self.authenticator.debug_authentication_state()}"
)
logger.debug(f"Auth state: {self.authenticator.debug_authentication_state()}")
raise Exception(f"No entitlement token in response for {content_id}: {e}")
except Exception as e:
logger.error(f"Error getting entitlement token for {content_id}: {e}")
logger.debug(
f"Auth state: {self.authenticator.debug_authentication_state()}"
)
logger.debug(f"Auth state: {self.authenticator.debug_authentication_state()}")
raise Exception(f"Error getting entitlement token for {content_id}: {e}")
def get_channel_playlist(self, channel_id: str, entitlement_token: str) -> Dict:
"""Get channel playlist data"""
if self.endpoint_manager and self.endpoint_manager.has_endpoint(
"channel_playlist"
):
if self.endpoint_manager and self.endpoint_manager.has_endpoint("channel_playlist"):
url = self.endpoint_manager.get_endpoint("channel_playlist").format(
channel_id=channel_id
)
@@ -1122,16 +1048,10 @@ class Magenta2Provider(StreamingProvider):
logger.debug(f"Getting playlist data for: {channel.name}")
playlist_data = self.get_channel_playlist(
channel.channel_id, entitlement_token
)
playlist_data = self.get_channel_playlist(channel.channel_id, entitlement_token)
manifest_url = playlist_data.get(
"manifestUrl", playlist_data.get("manifest")
)
license_url = playlist_data.get(
"licenseUrl", playlist_data.get("license")
)
manifest_url = playlist_data.get("manifestUrl", playlist_data.get("manifest"))
license_url = playlist_data.get("licenseUrl", playlist_data.get("license"))
certificate_url = playlist_data.get(
"certificateUrl", playlist_data.get("certificate")
)
@@ -1160,14 +1080,10 @@ class Magenta2Provider(StreamingProvider):
except Exception as e:
retries += 1
if retries < max_retries:
logger.debug(
f"Retry {retries}/{max_retries} for {channel.name}: {e}"
)
logger.debug(f"Retry {retries}/{max_retries} for {channel.name}: {e}")
time.sleep(1)
else:
logger.error(
f"Failed to get streaming data for {channel.name}: {e}"
)
logger.error(f"Failed to get streaming data for {channel.name}: {e}")
logger.info(f"Streaming data population complete:")
logger.info(f" Successful: {len(successful_channels)}")
@@ -1187,13 +1103,9 @@ class Magenta2Provider(StreamingProvider):
content_id=channel.channel_id, content_type=channel.content_type
)
playlist_data = self.get_channel_playlist(
channel.channel_id, entitlement_token
)
playlist_data = self.get_channel_playlist(channel.channel_id, entitlement_token)
manifest_url = playlist_data.get(
"manifestUrl", playlist_data.get("manifest")
)
manifest_url = playlist_data.get("manifestUrl", playlist_data.get("manifest"))
if not manifest_url:
return None
@@ -1256,9 +1168,7 @@ class Magenta2Provider(StreamingProvider):
)
# First check for error cases
error_title_pattern = (
r'<ref[^>]*title="([^"]*)"[^>]*abstract="([^"]*)"[^>]*>'
)
error_title_pattern = r'<ref[^>]*title="([^"]*)"[^>]*abstract="([^"]*)"[^>]*>'
error_match = re.search(error_title_pattern, smil_content)
if error_match:
@@ -1270,14 +1180,10 @@ class Magenta2Provider(StreamingProvider):
# Check for specific error patterns
if "errorFiles/Unavailable.flv" in smil_content:
logger.error(
f"SMIL returned unavailable content for channel {channel_id}"
)
logger.error(f"SMIL returned unavailable content for channel {channel_id}")
return None
if "Invalid Token" in title or "InvalidAuthToken" in smil_content:
logger.error(
f"Invalid authentication token for channel {channel_id}"
)
logger.error(f"Invalid authentication token for channel {channel_id}")
return None
if "403" in smil_content:
logger.error(f"Access forbidden (403) for channel {channel_id}")
@@ -1325,9 +1231,7 @@ class Magenta2Provider(StreamingProvider):
return mpd_url
# If we get here, no MPD URL was found
logger.warning(
f"No MPD URL found in SMIL response for channel {channel_id}"
)
logger.warning(f"No MPD URL found in SMIL response for channel {channel_id}")
# Log the full SMIL content for debugging in case of unexpected format
if len(smil_content) < 1000: # Only log if it's reasonably short
@@ -1380,26 +1284,18 @@ class Magenta2Provider(StreamingProvider):
start_iso = TimestampConverter.epoch_to_iso(
start_time, format_type="basic", as_utc=True
)
end_iso = TimestampConverter.epoch_to_iso(
end_time, format_type="basic", as_utc=True
)
end_iso = TimestampConverter.epoch_to_iso(end_time, format_type="basic", as_utc=True)
# Build the catchup manifest URL
# Check if the manifest already has query parameters
separator = "&" if "?" in base_manifest else "?"
catchup_manifest = (
f"{base_manifest}{separator}begin={start_iso}&end={end_iso}"
)
catchup_manifest = f"{base_manifest}{separator}begin={start_iso}&end={end_iso}"
logger.debug(
f"Catchup manifest for channel {channel_id}: {catchup_manifest}"
)
logger.debug(f"Catchup manifest for channel {channel_id}: {catchup_manifest}")
return catchup_manifest
except Exception as e:
logger.error(
f"Error building catchup manifest for channel {channel_id}: {e}"
)
logger.error(f"Error building catchup manifest for channel {channel_id}: {e}")
# Fall back to live manifest if catchup formatting fails
logger.warning(f"Falling back to live manifest for channel {channel_id}")
return base_manifest
@@ -1449,9 +1345,7 @@ class Magenta2Provider(StreamingProvider):
self._smil_cache[channel_id] = (now, smil_data)
logger.debug(f"Cached SMIL data for {channel_id}")
else:
logger.warning(
f"No MPD URL or releasePid found for {channel_id}, not caching"
)
logger.warning(f"No MPD URL or releasePid found for {channel_id}, not caching")
return smil_data
@@ -1466,9 +1360,7 @@ class Magenta2Provider(StreamingProvider):
try:
logger.debug("🔵 Step 1: Calling _ensure_authenticated()")
persona_token = self._ensure_authenticated()
logger.debug(
f"🔵 _ensure_authenticated() SUCCESS, token length: {len(persona_token)}"
)
logger.debug(f"🔵 _ensure_authenticated() SUCCESS, token length: {len(persona_token)}")
if not persona_token:
logger.error(f"No persona token!:")
@@ -1499,9 +1391,7 @@ class Magenta2Provider(StreamingProvider):
for key, value in headers.items():
if key == "Authorization":
# Don't log full auth token for security, but show it exists
logger.debug(
f" {key}: Basic [REDACTED] (length: {len(persona_token)})"
)
logger.debug(f" {key}: Basic [REDACTED] (length: {len(persona_token)})")
else:
logger.debug(f" {key}: {value}")
@@ -1534,16 +1424,12 @@ class Magenta2Provider(StreamingProvider):
smil_content,
self.http_manager,
client_id=client_id, # Use the same client_id as SMIL request
user_agent=self.platform_config[
"user_agent"
], # Platform user agent
user_agent=self.platform_config["user_agent"], # Platform user agent
)
return smil_content
else:
logger.error(
f"Failed to get SMIL content for DRM: {response.status_code}"
)
logger.error(f"Failed to get SMIL content for DRM: {response.status_code}")
return None
except Exception as e:
@@ -1630,9 +1516,7 @@ class Magenta2Provider(StreamingProvider):
# Debug what we do have
logger.debug(f"SMIL data keys: {list(smil_data.keys())}")
if "content" in smil_data and smil_data["content"]:
logger.debug(
f"SMIL content preview: {smil_data['content'][:500]}..."
)
logger.debug(f"SMIL content preview: {smil_data['content'][:500]}...")
return []
release_pid = smil_data["release_pid"]
@@ -1807,9 +1691,7 @@ class Magenta2Provider(StreamingProvider):
if current_time >= (expires_at - 300):
# Token expired but might be refreshable
# Check if we have yo_digital token with refresh capability
yo_token = context.get_token(
self.provider_name, "yo_digital", self.country
)
yo_token = context.get_token(self.provider_name, "yo_digital", self.country)
if yo_token and "refresh_token" in yo_token:
# Check if yo_digital refresh token is still valid
if (
@@ -1860,40 +1742,26 @@ class Magenta2Provider(StreamingProvider):
expires_at = token["expires_at"]
scope_info["expires_at"] = expires_at
scope_info["is_expired"] = current_time >= expires_at
scope_info["time_remaining"] = int(
max(0, expires_at - current_time)
)
scope_info["time_remaining"] = int(max(0, expires_at - current_time))
if "composed_at" in token:
scope_info["composed_at"] = token["composed_at"]
# Handle yo_digital token
elif scope == "yo_digital":
if (
"access_token_expires_in" in token
and "access_token_issued_at" in token
):
if "access_token_expires_in" in token and "access_token_issued_at" in token:
current_time = time.time()
expires_at = (
token["access_token_issued_at"]
+ token["access_token_expires_in"]
)
expires_at = token["access_token_issued_at"] + token["access_token_expires_in"]
scope_info["access_token_expires_at"] = expires_at
scope_info["access_token_is_expired"] = current_time >= expires_at
if (
"refresh_token_expires_in" in token
and "refresh_token_issued_at" in token
):
if "refresh_token_expires_in" in token and "refresh_token_issued_at" in token:
current_time = time.time()
refresh_expires_at = (
token["refresh_token_issued_at"]
+ token["refresh_token_expires_in"]
token["refresh_token_issued_at"] + token["refresh_token_expires_in"]
)
scope_info["refresh_token_expires_at"] = refresh_expires_at
scope_info["refresh_token_is_expired"] = (
current_time >= refresh_expires_at
)
scope_info["refresh_token_is_expired"] = current_time >= refresh_expires_at
scope_info["has_refresh_token"] = True
# Standard token handling (tvhubs, taa)
@@ -1973,9 +1841,7 @@ class Magenta2Provider(StreamingProvider):
result["endpoints"] = {
"has_taa_auth": self.endpoint_manager.has_endpoint("taa_auth"),
"has_entitlement": self.endpoint_manager.has_endpoint("entitlement"),
"has_widevine_license": self.endpoint_manager.has_endpoint(
"widevine_license"
),
"has_widevine_license": self.endpoint_manager.has_endpoint("widevine_license"),
"has_mpx_selector": self.endpoint_manager.has_endpoint("mpx_selector"),
"total_endpoints": len(self.endpoint_manager.get_all_endpoints()),
}
@@ -12,11 +12,9 @@ from typing import Any, Dict, Optional
from urllib.parse import parse_qs, unquote, urlparse
from ...base.network import HTTPManager
from ...base.ui import (NotificationFactory, NotificationInterface,
NotificationResult)
from ...base.ui import NotificationFactory, NotificationInterface, NotificationResult
from ...base.utils.logger import logger
from .constants import (DEFAULT_PLATFORM, DEFAULT_REQUEST_TIMEOUT, GRANT_TYPES,
MAGENTA2_PLATFORMS)
from .constants import DEFAULT_PLATFORM, DEFAULT_REQUEST_TIMEOUT, GRANT_TYPES, MAGENTA2_PLATFORMS
@dataclass
@@ -82,18 +80,14 @@ class RemoteLoginHandler:
self._current_session: Optional[RemoteLoginSession] = None
logger.debug(
f"RemoteLoginHandler initialized with {self._notifier.__class__.__name__}"
)
logger.debug(f"RemoteLoginHandler initialized with {self._notifier.__class__.__name__}")
def set_notifier(self, notifier: NotificationInterface) -> None:
"""Set custom notification interface"""
self._notifier = notifier
logger.debug(f"Notifier set to: {notifier.__class__.__name__}")
def start_remote_login(
self, scope: str = "tvhubs offline_access"
) -> RemoteLoginSession:
def start_remote_login(self, scope: str = "tvhubs offline_access") -> RemoteLoginSession:
"""
Start backchannel authentication flow
@@ -306,9 +300,7 @@ class RemoteLoginHandler:
# Perform poll
poll_count += 1
logger.debug(
f"Poll {poll_count}/{max_polls} (remaining: {remaining:.0f}s)"
)
logger.debug(f"Poll {poll_count}/{max_polls} (remaining: {remaining:.0f}s)")
try:
response = self.http_manager.post(
@@ -99,9 +99,7 @@ class Sam3Client:
def _get_token_endpoint(self) -> str:
"""Get the appropriate token endpoint with fallback logic"""
return (
self.oauth_token_endpoint or self.token_endpoint or self.line_auth_endpoint
)
return self.oauth_token_endpoint or self.token_endpoint or self.line_auth_endpoint
def _get_line_auth_endpoint(self) -> str:
"""Get line auth endpoint with fallback"""
@@ -277,9 +275,7 @@ class Sam3Client:
code, state = self._extract_code_and_state(redirect_url)
if not code or not state:
raise Exception(
"Could not extract authorization code and state from redirect"
)
raise Exception("Could not extract authorization code and state from redirect")
logger.info("SAM3 login completed successfully")
return {"code": code, "state": state, "redirect_url": redirect_url}
@@ -324,9 +320,7 @@ class Sam3Client:
logger.info("Line authentication successful, refresh token obtained")
return True
logger.warning(
"Line authentication succeeded but no refresh token received"
)
logger.warning("Line authentication succeeded but no refresh token received")
return False
except Exception as e:
@@ -403,9 +397,7 @@ class Sam3Client:
"""Check if remote login is available"""
return self._get_remote_login_handler() is not None
def remote_login(
self, scope: str = "tvhubs offline_access"
) -> Optional[Dict[str, Any]]:
def remote_login(self, scope: str = "tvhubs offline_access") -> Optional[Dict[str, Any]]:
"""
Perform complete remote login (backchannel auth) flow
@@ -505,9 +497,7 @@ class Sam3Client:
logger.warning(f"Token refresh failed for scope {scope}: {e}")
# If no refresh token available, we can't get an access token
logger.warning(
f"No access token available for scope {scope} - line auth may be needed"
)
logger.warning(f"No access token available for scope {scope} - line auth may be needed")
return ""
def get_token_endpoint(self) -> str:
@@ -575,9 +565,7 @@ class Sam3Client:
form_html = html_content[form_start:form_end]
# Find all hidden input fields
pattern = (
r'<input[^>]*type="hidden"[^>]*name="([^"]*)"[^>]*value="([^"]*)"[^>]*>'
)
pattern = r'<input[^>]*type="hidden"[^>]*name="([^"]*)"[^>]*value="([^"]*)"[^>]*>'
matches = re.findall(pattern, form_html, re.IGNORECASE)
for name, value in matches:
@@ -688,9 +676,7 @@ class Sam3Client:
# Get backchannel auth start endpoint
if "backchannel_auth_start" in openid_config:
self.backchannel_start_url = openid_config["backchannel_auth_start"]
logger.debug(
f"Backchannel auth endpoint from OpenID: {self.backchannel_start_url}"
)
logger.debug(f"Backchannel auth endpoint from OpenID: {self.backchannel_start_url}")
# OAuth token endpoint might be different from line auth endpoint
if "token_endpoint" in openid_config:
@@ -95,9 +95,7 @@ class TaaClient:
)
logger.debug(f"yo_digital request to: {endpoint}")
logger.debug(
f"TAA payload keyValue: {taa_payload.get('keyValue', 'MISSING')}"
)
logger.debug(f"TAA payload keyValue: {taa_payload.get('keyValue', 'MISSING')}")
# Perform yo_digital request
response = self.http_manager.post(
@@ -119,12 +117,8 @@ class TaaClient:
)
# Check for deviceLimitExceed (note: might be "Exceed" not "Exceeded")
if error_data.get("deviceLimitExceed") or error_data.get(
"deviceLimitExceeded"
):
logger.error(
"Device limit exceeded in yo_digital authentication"
)
if error_data.get("deviceLimitExceed") or error_data.get("deviceLimitExceeded"):
logger.error("Device limit exceeded in yo_digital authentication")
return YoDigitalTokens(
access_token="",
access_token_expires_in=0,
@@ -135,17 +129,13 @@ class TaaClient:
# Log any other error details
if "error" in error_data:
logger.error(
f"yo_digital error type: {error_data.get('error')}"
)
logger.error(f"yo_digital error type: {error_data.get('error')}")
if "error_description" in error_data:
logger.error(
f"yo_digital error description: {error_data.get('error_description')}"
)
if "message" in error_data:
logger.error(
f"yo_digital error message: {error_data.get('message')}"
)
logger.error(f"yo_digital error message: {error_data.get('message')}")
except (ValueError, KeyError) as e:
logger.error(f"Could not parse 400 error response: {e}")
@@ -314,9 +304,7 @@ class TaaClient:
error_data = response.json()
if error_data.get("deviceLimitExceeded"):
logger.error("Device limit exceeded in TAA authentication")
return TaaAuthResult(
access_token="", device_limit_exceeded=True
)
return TaaAuthResult(access_token="", device_limit_exceeded=True)
except:
pass
@@ -399,9 +387,7 @@ class TaaClient:
)
# Get TAA-specific OS format
resolved_os = (
self.platform_config.get("taa_os") or self.platform_config["firmware"]
)
resolved_os = self.platform_config.get("taa_os") or self.platform_config["firmware"]
resolved_client_model = client_model or f"ftv-{self.platform}"
@@ -110,9 +110,7 @@ class TokenFlowManager:
if not token_result.success or not token_result.access_token:
logger.debug("=== GET_PERSONA_TOKEN FAILED (token_result failed) ===")
return PersonaResult(
success=False, error=token_result.error or "No access token"
)
return PersonaResult(success=False, error=token_result.error or "No access token")
# Compose persona token with expiry information using existing method
from .token_utils import PersonaTokenComposer
@@ -161,9 +159,7 @@ class TokenFlowManager:
current_time = time.time()
expires_at = persona_data["expires_at"]
logger.debug(
f"🟡 Current time: {current_time}, Expires at: {expires_at}"
)
logger.debug(f"🟡 Current time: {current_time}, Expires at: {expires_at}")
# Check if cached token is still valid using the actual persona JWT expiry
if current_time < (expires_at - 300): # 5-minute buffer
@@ -176,9 +172,7 @@ class TokenFlowManager:
expires_at=expires_at, # 🆕 Return expiry
)
else:
logger.debug(
f"🟡 Cached persona token expired at {time.ctime(expires_at)}"
)
logger.debug(f"🟡 Cached persona token expired at {time.ctime(expires_at)}")
else:
logger.debug("🟡 No valid persona data in cache")
@@ -214,9 +208,7 @@ class TokenFlowManager:
"""Backward compatibility method - delegates to new composition"""
from .token_utils import PersonaTokenComposer
result = PersonaTokenComposer.compose_from_jwt(
access_token, MAGENTA2_FALLBACK_ACCOUNT_URI
)
result = PersonaTokenComposer.compose_from_jwt(access_token, MAGENTA2_FALLBACK_ACCOUNT_URI)
return result.persona_token if result else None
def get_yo_digital_token(self, force_refresh: bool = False) -> TokenFlowResult:
@@ -312,9 +304,7 @@ class TokenFlowManager:
except Exception as e:
logger.debug(f"Error checking yo_digital access_token: {e}")
return TokenFlowResult(
success=False, error=str(e), flow_path="check_yo_digital_access"
)
return TokenFlowResult(success=False, error=str(e), flow_path="check_yo_digital_access")
# ========================================================================
# Step 2: Refresh yo_digital tokens
@@ -344,9 +334,7 @@ class TokenFlowManager:
# Refresh via TaaClient (stub for now)
logger.debug("Attempting to refresh yo_digital tokens")
new_tokens_dict = self.taa_client.refresh_yo_digital_tokens(
token_data["refresh_token"]
)
new_tokens_dict = self.taa_client.refresh_yo_digital_tokens(token_data["refresh_token"])
if not new_tokens_dict:
return TokenFlowResult(
@@ -369,9 +357,7 @@ class TokenFlowManager:
except Exception as e:
logger.debug(f"Error refreshing yo_digital tokens: {e}")
return TokenFlowResult(
success=False, error=str(e), flow_path="refresh_yo_digital"
)
return TokenFlowResult(success=False, error=str(e), flow_path="refresh_yo_digital")
# ========================================================================
# Step 3: Get yo_digital from taa access_token
@@ -404,9 +390,7 @@ class TokenFlowManager:
logger.debug("Getting yo_digital tokens from taa access_token")
yo_digital_result = self.taa_client.get_yo_digital_tokens(
taa_access_token=taa_token_data["access_token"],
device_id=self.session_manager.get_device_id(
self.provider_name, self.country
),
device_id=self.session_manager.get_device_id(self.provider_name, self.country),
)
if not yo_digital_result:
@@ -419,12 +403,8 @@ class TokenFlowManager:
logger.info("Clearing cached tokens and falling back to remote_login")
# Clear the invalid cached tokens to prevent repeated failures
self.session_manager.clear_scoped_token(
self.provider_name, "taa", self.country
)
self.session_manager.clear_scoped_token(
self.provider_name, "tvhubs", self.country
)
self.session_manager.clear_scoped_token(self.provider_name, "taa", self.country)
self.session_manager.clear_scoped_token(self.provider_name, "tvhubs", self.country)
# Fallback to remote_login
return self._get_yo_digital_via_remote_login()
@@ -449,9 +429,7 @@ class TokenFlowManager:
except Exception as e:
logger.debug(f"Error getting yo_digital from taa: {e}")
return TokenFlowResult(
success=False, error=str(e), flow_path="yo_digital_from_taa"
)
return TokenFlowResult(success=False, error=str(e), flow_path="yo_digital_from_taa")
# ========================================================================
# Step 4: Exchange shared refresh_token for taa, then yo_digital
@@ -461,9 +439,7 @@ class TokenFlowManager:
"""Exchange shared refresh_token for taa, then get yo_digital"""
try:
# Check if we have shared refresh_token at provider level
session_data = self.session_manager.load_session(
self.provider_name, self.country
)
session_data = self.session_manager.load_session(self.provider_name, self.country)
if not session_data or "refresh_token" not in session_data:
return TokenFlowResult(
@@ -494,9 +470,7 @@ class TokenFlowManager:
logger.debug("Getting yo_digital tokens from exchanged taa token")
yo_digital_result = self.taa_client.get_yo_digital_tokens(
taa_access_token=taa_token,
device_id=self.session_manager.get_device_id(
self.provider_name, self.country
),
device_id=self.session_manager.get_device_id(self.provider_name, self.country),
)
if not yo_digital_result:
@@ -509,17 +483,11 @@ class TokenFlowManager:
logger.info("Clearing cached tokens and falling back to remote_login")
# Clear the invalid cached tokens and refresh_token
self.session_manager.clear_scoped_token(
self.provider_name, "taa", self.country
)
self.session_manager.clear_scoped_token(
self.provider_name, "tvhubs", self.country
)
self.session_manager.clear_scoped_token(self.provider_name, "taa", self.country)
self.session_manager.clear_scoped_token(self.provider_name, "tvhubs", self.country)
# Clear the shared refresh_token from session data
session_data = self.session_manager.load_session(
self.provider_name, self.country
)
session_data = self.session_manager.load_session(self.provider_name, self.country)
if session_data and "refresh_token" in session_data:
del session_data["refresh_token"]
self.session_manager.save_session(
@@ -609,9 +577,7 @@ class TokenFlowManager:
logger.debug("Getting yo_digital tokens from line_auth taa token")
yo_digital_result = self.taa_client.get_yo_digital_tokens(
taa_access_token=taa_token,
device_id=self.session_manager.get_device_id(
self.provider_name, self.country
),
device_id=self.session_manager.get_device_id(self.provider_name, self.country),
)
if not yo_digital_result:
@@ -621,9 +587,7 @@ class TokenFlowManager:
"yo_digital acquisition failed after successful line_auth - "
"likely ISP account differs from TV subscription account"
)
logger.info(
"Falling back to remote_login for TV account authentication"
)
logger.info("Falling back to remote_login for TV account authentication")
# Fallback to remote_login (tokens will be overwritten)
return self._get_yo_digital_via_remote_login()
@@ -677,9 +641,7 @@ class TokenFlowManager:
logger.info("Attempting remote_login flow")
# Perform remote login
remote_token_data = self.sam3_client.remote_login(
scope="tvhubs offline_access"
)
remote_token_data = self.sam3_client.remote_login(scope="tvhubs offline_access")
if not remote_token_data:
return TokenFlowResult(
@@ -717,9 +679,7 @@ class TokenFlowManager:
logger.debug("Getting yo_digital tokens from remote_login taa token")
yo_digital_result = self.taa_client.get_yo_digital_tokens(
taa_access_token=taa_token,
device_id=self.session_manager.get_device_id(
self.provider_name, self.country
),
device_id=self.session_manager.get_device_id(self.provider_name, self.country),
)
if not yo_digital_result:
@@ -766,9 +726,7 @@ class TokenFlowManager:
):
return False
expires_at = (
token_data["access_token_issued_at"] + token_data["access_token_expires_in"]
)
expires_at = token_data["access_token_issued_at"] + token_data["access_token_expires_in"]
# Use 5 minute buffer
return time.time() < (expires_at - 300)
@@ -781,10 +739,7 @@ class TokenFlowManager:
):
return False
expires_at = (
token_data["refresh_token_issued_at"]
+ token_data["refresh_token_expires_in"]
)
expires_at = token_data["refresh_token_issued_at"] + token_data["refresh_token_expires_in"]
# Use 5 minute buffer
return time.time() < (expires_at - 300)
@@ -836,9 +791,7 @@ class TokenFlowManager:
"issued_at": time.time(),
}
self.session_manager.save_scoped_token(
self.provider_name, "taa", token_data, self.country
)
self.session_manager.save_scoped_token(self.provider_name, "taa", token_data, self.country)
logger.debug("taa token saved")
def _save_tvhubs_token(self, token_data: Dict[str, Any]) -> None:
@@ -857,18 +810,14 @@ class TokenFlowManager:
def _save_refresh_token(self, refresh_token: str) -> None:
"""Save shared refresh_token at provider level"""
session_data = (
self.session_manager.load_session(self.provider_name, self.country) or {}
)
session_data = self.session_manager.load_session(self.provider_name, self.country) or {}
session_data["refresh_token"] = refresh_token
session_data["device_id"] = self.session_manager.get_device_id(
self.provider_name, self.country
)
self.session_manager.save_session(
self.provider_name, session_data, self.country
)
self.session_manager.save_session(self.provider_name, session_data, self.country)
logger.debug("Shared refresh_token saved")
# ========================================================================
@@ -905,9 +854,7 @@ class TokenFlowManager:
def _get_taa_status(self) -> Dict[str, Any]:
"""Get taa token status"""
token_data = self.session_manager.load_scoped_token(
self.provider_name, "taa", self.country
)
token_data = self.session_manager.load_scoped_token(self.provider_name, "taa", self.country)
if not token_data:
return {"exists": False}
@@ -935,9 +882,7 @@ class TokenFlowManager:
def _get_refresh_token_status(self) -> Dict[str, Any]:
"""Get shared refresh_token status"""
session_data = self.session_manager.load_session(
self.provider_name, self.country
)
session_data = self.session_manager.load_session(self.provider_name, self.country)
if not session_data or "refresh_token" not in session_data:
return {"exists": False}
@@ -46,9 +46,7 @@ class JWTClaims:
def is_user_token(self) -> bool:
"""Check if this is a user-authenticated token"""
return bool(
self.persona_id or self.account_id or self.consumer_id or self.tv_account_id
)
return bool(self.persona_id or self.account_id or self.consumer_id or self.tv_account_id)
class JWTParser:
@@ -203,9 +201,7 @@ class PersonaTokenComposer:
"""Compose persona token and return with expiry information"""
logger.debug(f"🟢 PersonaTokenComposer.compose_from_jwt START")
try:
logger.debug(
f"🟢 Input JWT token length: {len(jwt_token) if jwt_token else 0}"
)
logger.debug(f"🟢 Input JWT token length: {len(jwt_token) if jwt_token else 0}")
claims = JWTParser.parse(jwt_token)
if not claims:
@@ -225,9 +221,7 @@ class PersonaTokenComposer:
return None
# Compose the persona token
composed_token = PersonaTokenComposer._compose_token(
account_uri, persona_jwt
)
composed_token = PersonaTokenComposer._compose_token(account_uri, persona_jwt)
if not composed_token:
logger.error("🔴 _compose_token returned None")
return None
@@ -279,9 +273,7 @@ class PersonaTokenComposer:
return None
@staticmethod
def compose_from_components(
account_uri: str, dc_cts_persona_token: str
) -> Optional[str]:
def compose_from_components(account_uri: str, dc_cts_persona_token: str) -> Optional[str]:
"""Original method - for backward compatibility"""
return PersonaTokenComposer._compose_token(account_uri, dc_cts_persona_token)
@@ -315,9 +307,7 @@ class PersonaTokenComposer:
# Verify persona_jwt looks like a JWT
if not persona_jwt.startswith("eyJ"):
logger.warning(
f"Extracted token doesn't look like a JWT: {persona_jwt[:20]}..."
)
logger.warning(f"Extracted token doesn't look like a JWT: {persona_jwt[:20]}...")
return {"account_uri": account_uri, "persona_jwt": persona_jwt}
@@ -1,7 +1,6 @@
# streaming_providers/providers/magenta_eu/__init__.py
from .auth import MagentaAuthenticator, MagentaAuthToken
from .constants import (API_ENDPOINTS, COUNTRY_CONFIG, DEFAULT_COUNTRY,
SUPPORTED_COUNTRIES)
from .constants import API_ENDPOINTS, COUNTRY_CONFIG, DEFAULT_COUNTRY, SUPPORTED_COUNTRIES
from .provider import MagentaEUProvider
__all__ = [
@@ -17,19 +17,37 @@ except ImportError:
from Crypto.Cipher import PKCS1_OAEP
from Crypto.PublicKey import RSA
from ...base.auth.base_auth import (BaseAuthenticator, BaseAuthToken,
TokenAuthLevel)
from ...base.auth.base_auth import BaseAuthenticator, BaseAuthToken, TokenAuthLevel
from ...base.models.proxy_models import ProxyConfig
from ...base.utils.logger import logger
from .constants import (API_ENDPOINTS, APP_VERSION, AUTH_FLOWS, AUTH_STEPS,
BROADCASTING_STREAM_LIMITATION_APPLIES, CALL_TYPES,
CHANNEL_ID, COUNTRY_CONFIG, DEFAULT_COUNTRY,
DEFAULT_REQUEST_TIMEOUT, DEVICE_CONCURRENCY_PARAM,
DEVICE_MANUFACTURER, DEVICE_MODEL, DEVICE_NAME,
DEVICE_OS, DEVICE_TYPE, LOGIN_CONTEXT, LOGIN_TYPE,
MANAGE_DEVICE, SUPPORTED_COUNTRIES, USER_AGENT,
X_USER_AGENT, get_base_headers, get_base_url,
get_bifrost_url, get_language)
from .constants import (
API_ENDPOINTS,
APP_VERSION,
AUTH_FLOWS,
AUTH_STEPS,
BROADCASTING_STREAM_LIMITATION_APPLIES,
CALL_TYPES,
CHANNEL_ID,
COUNTRY_CONFIG,
DEFAULT_COUNTRY,
DEFAULT_REQUEST_TIMEOUT,
DEVICE_CONCURRENCY_PARAM,
DEVICE_MANUFACTURER,
DEVICE_MODEL,
DEVICE_NAME,
DEVICE_OS,
DEVICE_TYPE,
LOGIN_CONTEXT,
LOGIN_TYPE,
MANAGE_DEVICE,
SUPPORTED_COUNTRIES,
USER_AGENT,
X_USER_AGENT,
get_base_headers,
get_base_url,
get_bifrost_url,
get_language,
)
class InvalidTokenError(Exception):
@@ -89,9 +107,7 @@ class MagentaAuthToken(BaseAuthToken):
"expires_in": self.expires_in,
"issued_at": self.issued_at,
"auth_level": (
self.auth_level.value
if self.auth_level
else TokenAuthLevel.UNKNOWN.value
self.auth_level.value if self.auth_level else TokenAuthLevel.UNKNOWN.value
),
"credential_type": self.credential_type or "",
}
@@ -164,9 +180,7 @@ class MagentaAuthConfig:
try:
rsa_key = self.country_config["rsa_key"]
if not rsa_key:
logger.error(
f"No RSA public key configured for country: {self.country}"
)
logger.error(f"No RSA public key configured for country: {self.country}")
return password
key = RSA.import_key(rsa_key)
@@ -265,9 +279,7 @@ class MagentaAuthenticator(BaseAuthenticator):
logger.warning(f"CRITICAL: No session IDs available, using random fallback")
# Ensure current token has the correct IDs
if not self._current_token or not isinstance(
self._current_token, MagentaAuthToken
):
if not self._current_token or not isinstance(self._current_token, MagentaAuthToken):
self._current_token = MagentaAuthToken(
access_token="",
refresh_token="",
@@ -370,9 +382,7 @@ class MagentaAuthenticator(BaseAuthenticator):
"""Build authentication payload - required by BaseAuthenticator"""
from ...base.auth.credentials import UserPasswordCredentials
if not self.credentials or not isinstance(
self.credentials, UserPasswordCredentials
):
if not self.credentials or not isinstance(self.credentials, UserPasswordCredentials):
raise Exception("No valid credentials available")
# Enhanced validation
@@ -416,9 +426,7 @@ class MagentaAuthenticator(BaseAuthenticator):
},
}
def _create_token_from_response(
self, response_data: Dict[str, Any]
) -> BaseAuthToken:
def _create_token_from_response(self, response_data: Dict[str, Any]) -> BaseAuthToken:
"""Create token from API response - required by BaseAuthenticator"""
# PRESERVE the existing session IDs (which follow the correct priority)
device_id = ""
@@ -455,28 +463,20 @@ class MagentaAuthenticator(BaseAuthenticator):
# DUAL KEY SUPPORT: Handle both camelCase (API responses) and snake_case (stored sessions)
# Access token
access_token = response_data.get("accessToken") or response_data.get(
"access_token"
)
access_token = response_data.get("accessToken") or response_data.get("access_token")
if not access_token:
logger.error(f"CRITICAL: No access token found in response data")
logger.error(f"Available keys: {list(response_data.keys())}")
raise Exception("No access token found in response data")
# Refresh token
refresh_token = response_data.get("refreshToken") or response_data.get(
"refresh_token", ""
)
refresh_token = response_data.get("refreshToken") or response_data.get("refresh_token", "")
# Expires in
expires_in = response_data.get("expiresIn") or response_data.get(
"expires_in", 3600
)
expires_in = response_data.get("expiresIn") or response_data.get("expires_in", 3600)
# Token type
token_type = response_data.get("tokenType") or response_data.get(
"token_type", "Bearer"
)
token_type = response_data.get("tokenType") or response_data.get("token_type", "Bearer")
# For stored sessions, issued_at might be in the data, otherwise use current time
issued_at = response_data.get("issued_at", time.time())
@@ -494,9 +494,7 @@ class MagentaAuthenticator(BaseAuthenticator):
# Classify token
token.auth_level = self._classify_token(token)
logger.debug(
f"Token created successfully from {len(response_data)} data fields"
)
logger.debug(f"Token created successfully from {len(response_data)} data fields")
return token
@@ -531,9 +529,7 @@ class MagentaAuthenticator(BaseAuthenticator):
headers = self._get_auth_headers()
payload = self._build_auth_payload()
logger.debug(
f"Authentication payload prepared for user: {self.credentials.username}"
)
logger.debug(f"Authentication payload prepared for user: {self.credentials.username}")
response = self._http_manager.post(
self.auth_endpoint,
@@ -554,9 +550,7 @@ class MagentaAuthenticator(BaseAuthenticator):
return self._create_token_from_response(token_data)
except Exception as e:
logger.error(
f"Authentication failed for user {self.credentials.username}: {e}"
)
logger.error(f"Authentication failed for user {self.credentials.username}: {e}")
raise
def _upgrade_token(self, refresh_token: str) -> Dict[str, Any]:
@@ -648,9 +642,7 @@ class MagentaAuthenticator(BaseAuthenticator):
# Create new token with updated data but preserve session IDs
new_token = MagentaAuthToken(
access_token=token_data["accessToken"],
refresh_token=token_data.get(
"refreshToken", self._current_token.refresh_token
),
refresh_token=token_data.get("refreshToken", self._current_token.refresh_token),
token_type="Bearer",
expires_in=token_data.get("expiresIn", 3600),
issued_at=time.time(),
@@ -91,7 +91,9 @@ DEVICE_CONCURRENCY_PARAM = "TVSOA-restriction-unmanagedDeviceStreamLimit"
# User agent configuration
USER_AGENT = f"Mozilla/5.0 (X11; {OS} x86_64) AppleWebKit/537.36 (KHTML, like Gecko) {BROWSER}/{BROWSER_VERSION}.0.0.0 Safari/537.36"
X_USER_AGENT = f"{DEVICE_MODEL.lower()}|{DEVICE_TYPE.lower()}|{BROWSER}-{BROWSER_VERSION}|{APP_VERSION}|1"
X_USER_AGENT = (
f"{DEVICE_MODEL.lower()}|{DEVICE_TYPE.lower()}|{BROWSER}-{BROWSER_VERSION}|{APP_VERSION}|1"
)
# ============================================================================
# API Endpoints
@@ -11,13 +11,26 @@ from ...base.network import ProxyConfigManager
from ...base.provider import StreamingProvider
from ...base.utils.logger import logger
from .auth import MagentaAuthenticator
from .constants import (API_ENDPOINTS, CONTENT_TYPE_LIVE, DEFAULT_COUNTRY,
DEFAULT_MAX_RETRIES, DEFAULT_REQUEST_TIMEOUT,
DRM_SYSTEM_WIDEVINE, MAGENTA_TV_AT_LOGO,
MAGENTA_TV_PL_LOGO, MAX_TV_LOGO, STREAMING_FORMAT_DASH,
SUPPORTED_COUNTRIES, USER_AGENT, WV_URL, get_base_url,
get_bifrost_url, get_guest_headers, get_language,
get_natco_key)
from .constants import (
API_ENDPOINTS,
CONTENT_TYPE_LIVE,
DEFAULT_COUNTRY,
DEFAULT_MAX_RETRIES,
DEFAULT_REQUEST_TIMEOUT,
DRM_SYSTEM_WIDEVINE,
MAGENTA_TV_AT_LOGO,
MAGENTA_TV_PL_LOGO,
MAX_TV_LOGO,
STREAMING_FORMAT_DASH,
SUPPORTED_COUNTRIES,
USER_AGENT,
WV_URL,
get_base_url,
get_bifrost_url,
get_guest_headers,
get_language,
get_natco_key,
)
class MagentaEUProvider(StreamingProvider):
@@ -42,8 +55,7 @@ class MagentaEUProvider(StreamingProvider):
if not self.validate_country(country):
supported = ", ".join(self.SUPPORTED_COUNTRIES)
raise ValueError(
f"Unsupported country: {country}. "
f"MagentaTV EU supports: {supported}"
f"Unsupported country: {country}. " f"MagentaTV EU supports: {supported}"
)
super().__init__(country=country)
@@ -79,9 +91,7 @@ class MagentaEUProvider(StreamingProvider):
logger.info(f"=== MagentaProvider.__init__ COMPLETE ===")
def _load_proxy_from_manager(
self, config_dir: Optional[str]
) -> Optional[ProxyConfig]:
def _load_proxy_from_manager(self, config_dir: Optional[str]) -> Optional[ProxyConfig]:
"""Load proxy configuration from ProxyConfigManager"""
try:
proxy_manager = ProxyConfigManager(config_dir)
@@ -125,18 +135,14 @@ class MagentaEUProvider(StreamingProvider):
return ["user_credentials"]
def authenticate(self, **kwargs) -> str:
logger.info(
f"=== MagentaProvider.authenticate() CALLED with kwargs: {kwargs} ==="
)
logger.info(f"=== MagentaProvider.authenticate() CALLED with kwargs: {kwargs} ===")
self.bearer_token = self.authenticator.get_bearer_token(
force_refresh=kwargs.get("force_refresh", False)
)
logger.info(f"=== MagentaProvider.authenticate() COMPLETE ===")
return self.bearer_token
def get_dynamic_manifest_params(
self, channel: StreamingChannel, **kwargs
) -> Optional[str]:
def get_dynamic_manifest_params(self, channel: StreamingChannel, **kwargs) -> Optional[str]:
return None
def refresh_authentication(self) -> str:
@@ -173,9 +179,7 @@ class MagentaEUProvider(StreamingProvider):
"natco_code": self.country,
}
logger.debug(
f"Fetching channels with device_id: {device_id}, session_id: {session_id}"
)
logger.debug(f"Fetching channels with device_id: {device_id}, session_id: {session_id}")
response = self.http_manager.get(
channels_url,
@@ -192,9 +196,7 @@ class MagentaEUProvider(StreamingProvider):
self._channels_cache = channels
self._channels_cache_timestamp = time.time()
logger.info(
f"Successfully fetched {len(channels)} channels for country {self.country}"
)
logger.info(f"Successfully fetched {len(channels)} channels for country {self.country}")
return channels
except Exception as e:
@@ -230,9 +232,7 @@ class MagentaEUProvider(StreamingProvider):
if media_pid:
manifest_script_parts.append(f"media={media_pid}")
manifest_script = (
" ".join(manifest_script_parts) if manifest_script_parts else ""
)
manifest_script = " ".join(manifest_script_parts) if manifest_script_parts else ""
# Create streaming channel
streaming_channel = StreamingChannel(
@@ -327,35 +327,25 @@ class MagentaEUProvider(StreamingProvider):
start_iso = TimestampConverter.epoch_to_iso(
start_time, format_type="basic", as_utc=True
)
end_iso = TimestampConverter.epoch_to_iso(
end_time, format_type="basic", as_utc=True
)
end_iso = TimestampConverter.epoch_to_iso(end_time, format_type="basic", as_utc=True)
# Build the catchup manifest URL
# Check if the manifest already has query parameters
separator = "&" if "?" in base_manifest else "?"
catchup_manifest = (
f"{base_manifest}{separator}begin={start_iso}&end={end_iso}"
)
catchup_manifest = f"{base_manifest}{separator}begin={start_iso}&end={end_iso}"
logger.debug(
f"Catchup manifest for channel {channel_id}: {catchup_manifest}"
)
logger.debug(f"Catchup manifest for channel {channel_id}: {catchup_manifest}")
return catchup_manifest
except Exception as e:
logger.error(
f"Error building catchup manifest for channel {channel_id}: {e}"
)
logger.error(f"Error building catchup manifest for channel {channel_id}: {e}")
# Fall back to live manifest if catchup formatting fails
logger.warning(f"Falling back to live manifest for channel {channel_id}")
return base_manifest
def get_drm(self, channel_id: str, **kwargs) -> List[DRMConfig]:
"""Get DRM configurations for channel by ID"""
logger.info(
f"=== get_drm_configs_by_id CALLED for channel_id: {channel_id} ==="
)
logger.info(f"=== get_drm_configs_by_id CALLED for channel_id: {channel_id} ===")
# Find channel in cache
channel = None
@@ -379,9 +369,7 @@ class MagentaEUProvider(StreamingProvider):
drm_config = self.get_drm_config(channel)
return [drm_config] if drm_config else []
def get_drm_config(
self, channel: StreamingChannel, **kwargs
) -> Optional[DRMConfig]:
def get_drm_config(self, channel: StreamingChannel, **kwargs) -> Optional[DRMConfig]:
"""Get DRM configuration for channel with correct authentication"""
try:
import base64
@@ -450,9 +438,7 @@ class MagentaEUProvider(StreamingProvider):
# Build license URL with parameters
license_url = (
f"{WV_URL}{pid}&"
f"token={persona_token}&"
f"account={encoded_account_uri}"
f"{WV_URL}{pid}&" f"token={persona_token}&" f"account={encoded_account_uri}"
)
logger.debug(f"License URL created: {license_url[:100]}...")
@@ -8,8 +8,7 @@ from ...base.auth.base_oauth2_auth import BaseOAuth2Authenticator
from ...base.models.proxy_models import ProxyConfig
from ...base.utils.logger import logger
from .constants import RTLPlusConfig, RTLPlusDefaults
from .models import (RTLPlusAuthToken, RTLPlusClientCredentials,
RTLPlusUserCredentials)
from .models import RTLPlusAuthToken, RTLPlusClientCredentials, RTLPlusUserCredentials
class RTLPlusAuthenticator(BaseOAuth2Authenticator):
@@ -86,9 +85,7 @@ class RTLPlusAuthenticator(BaseOAuth2Authenticator):
config_creds = self._get_anonymous_credentials_from_config()
if config_creds:
return RTLPlusClientCredentials(
client_id=config_creds.get(
"client_id", RTLPlusDefaults.ANONYMOUS_CLIENT_ID
),
client_id=config_creds.get("client_id", RTLPlusDefaults.ANONYMOUS_CLIENT_ID),
client_secret=config_creds.get(
"client_secret", RTLPlusDefaults.ANONYMOUS_CLIENT_SECRET
),
@@ -99,9 +96,7 @@ class RTLPlusAuthenticator(BaseOAuth2Authenticator):
# Fallback to default credentials
return RTLPlusClientCredentials()
def _create_token_from_response(
self, response_data: Dict[str, Any]
) -> RTLPlusAuthToken:
def _create_token_from_response(self, response_data: Dict[str, Any]) -> RTLPlusAuthToken:
"""Create RTL+-specific token from OAuth2 response"""
import time
@@ -168,9 +163,7 @@ class RTLPlusAuthenticator(BaseOAuth2Authenticator):
# Check for user-authenticated token
if preferred_username or email:
logger.debug(
"RTL+ Token classified as USER_AUTHENTICATED (has user claims)"
)
logger.debug("RTL+ Token classified as USER_AUTHENTICATED (has user claims)")
return TokenAuthLevel.USER_AUTHENTICATED
# Check for client credentials (anonymous) token
@@ -321,9 +314,7 @@ class RTLPlusAuthenticator(BaseOAuth2Authenticator):
"""
from ...base.auth.credentials import UserPasswordCredentials
return isinstance(
self.credentials, (RTLPlusUserCredentials, UserPasswordCredentials)
)
return isinstance(self.credentials, (RTLPlusUserCredentials, UserPasswordCredentials))
def has_stored_credentials(self) -> bool:
"""
@@ -331,9 +322,7 @@ class RTLPlusAuthenticator(BaseOAuth2Authenticator):
"""
try:
logger.debug("RTL+ Checking for stored credentials using settings manager")
stored_creds = self.settings_manager.get_provider_credentials(
self.provider_name
)
stored_creds = self.settings_manager.get_provider_credentials(self.provider_name)
if not stored_creds:
logger.debug("RTL+ No stored credentials found")
@@ -366,9 +355,7 @@ class RTLPlusAuthenticator(BaseOAuth2Authenticator):
status.update(
{
"has_user_credentials": self.has_user_credentials(),
"authentication_mode": (
"user" if self.has_user_credentials() else "anonymous"
),
"authentication_mode": ("user" if self.has_user_credentials() else "anonymous"),
"client_version": self.config.client_version,
}
)
@@ -26,7 +26,9 @@ class RTLPlusDefaults:
AUTH_ENDPOINT = f"{AUTH_BASE_URL}/token"
AUTH_AUTHORIZE_ENDPOINT = f"{AUTH_BASE_URL}/auth"
GRAPHQL_ENDPOINT = "https://cdn.gateway.now-plus-prod.aws-cbc.cloud/graphql"
MANIFEST_ENDPOINT = "https://stus.player.streamingtech.de/livestream/linear/{channel_id}?platform=web"
MANIFEST_ENDPOINT = (
"https://stus.player.streamingtech.de/livestream/linear/{channel_id}?platform=web"
)
BASE_WEBSITE = "https://plus.rtl.de/"
CONFIG_ENDPOINT = "https://plus.rtl.de/assets/config/config.json"
@@ -88,8 +90,7 @@ class RTLPlusHeaders:
"Content-Type": "application/json",
"Rtlplus-Client-Id": RTLPlusDefaults.CLIENT_ID,
"Rtlplus-Referrer": "",
"Rtlplus-Client-Version": client_version
or RTLPlusDefaults.CLIENT_VERSION,
"Rtlplus-Client-Version": client_version or RTLPlusDefaults.CLIENT_VERSION,
}
)
@@ -102,9 +103,7 @@ class RTLPlusHeaders:
return headers
@staticmethod
def get_drm_headers(
access_token: str, device_id: str = None, user_agent: str = None
) -> dict:
def get_drm_headers(access_token: str, device_id: str = None, user_agent: str = None) -> dict:
"""Get headers for DRM license requests"""
return {
"X-Auth-Token": access_token,
@@ -126,27 +125,17 @@ class RTLPlusConfig:
self.logo = config.get("logo", RTLPlusDefaults.RTLPLUS_LOGO)
# Core settings (can be overridden)
self.client_version = config.get(
"client_version", RTLPlusDefaults.CLIENT_VERSION
)
self.chrome_version = config.get(
"chrome_version", RTLPlusDefaults.CHROME_VERSION
)
self.client_version = config.get("client_version", RTLPlusDefaults.CLIENT_VERSION)
self.chrome_version = config.get("chrome_version", RTLPlusDefaults.CHROME_VERSION)
self.device_id = config.get("device_id", RTLPlusDefaults.DEVICE_ID)
self.user_agent = config.get("user_agent", RTLPlusDefaults.USER_AGENT)
# API endpoints (can be overridden for testing)
self.auth_endpoint = config.get("auth_endpoint", RTLPlusDefaults.AUTH_ENDPOINT)
self.graphql_endpoint = config.get(
"graphql_endpoint", RTLPlusDefaults.GRAPHQL_ENDPOINT
)
self.manifest_endpoint = config.get(
"manifest_endpoint", RTLPlusDefaults.MANIFEST_ENDPOINT
)
self.graphql_endpoint = config.get("graphql_endpoint", RTLPlusDefaults.GRAPHQL_ENDPOINT)
self.manifest_endpoint = config.get("manifest_endpoint", RTLPlusDefaults.MANIFEST_ENDPOINT)
self.base_website = config.get("base_website", RTLPlusDefaults.BASE_WEBSITE)
self.config_endpoint = config.get(
"config_endpoint", RTLPlusDefaults.CONFIG_ENDPOINT
)
self.config_endpoint = config.get("config_endpoint", RTLPlusDefaults.CONFIG_ENDPOINT)
# HTTP settings
self.timeout = config.get("timeout", RTLPlusDefaults.DEFAULT_TIMEOUT)
@@ -40,9 +40,7 @@ class RTLPlusClientCredentials(ClientCredentials):
RTL+ specific client credentials (anonymous access)
"""
def __init__(
self, client_id: Optional[str] = None, client_secret: Optional[str] = None
):
def __init__(self, client_id: Optional[str] = None, client_secret: Optional[str] = None):
super().__init__(
client_id=client_id or RTLPlusDefaults.ANONYMOUS_CLIENT_ID,
client_secret=client_secret or RTLPlusDefaults.ANONYMOUS_CLIENT_SECRET,
@@ -200,9 +198,7 @@ class RTLPlusStreamInfo:
raise ValueError("Stream must have both manifest_url and channel_id")
@classmethod
def from_manifest_response(
cls, data: Dict[str, Any], channel_id: str
) -> "RTLPlusStreamInfo":
def from_manifest_response(cls, data: Dict[str, Any], channel_id: str) -> "RTLPlusStreamInfo":
"""Create stream info from manifest API response"""
return cls(
manifest_url=data["url"],
@@ -51,9 +51,7 @@ class RTLPlusProvider(StreamingProvider):
)
# ✅ Share HTTP manager with authenticator
self.http_manager = self._share_http_manager_with_authenticator(
self.authenticator
)
self.http_manager = self._share_http_manager_with_authenticator(self.authenticator)
# Try authentication
try:
@@ -259,9 +257,7 @@ class RTLPlusProvider(StreamingProvider):
logger.debug(f"RTL+ Manifest Request: GET {manifest_url}")
headers = self.rtl_config.get_base_headers()
response = self.http_manager.get(
manifest_url, operation="manifest", headers=headers
)
response = self.http_manager.get(manifest_url, operation="manifest", headers=headers)
logger.debug(f"RTL+ Manifest Response: Status={response.status_code}")
logger.debug(f"RTL+ Response Headers: {dict(response.headers)}")
@@ -269,9 +265,7 @@ class RTLPlusProvider(StreamingProvider):
response.raise_for_status()
manifest_data = response.json()
logger.debug(
f"RTL+ Manifest Data: {self._sanitize_manifest_log(manifest_data)}"
)
logger.debug(f"RTL+ Manifest Data: {self._sanitize_manifest_log(manifest_data)}")
# Process manifest data
quality_preference = ["dashhd", "dashsd"]
@@ -280,9 +274,7 @@ class RTLPlusProvider(StreamingProvider):
for stream in manifest_data:
if stream.get("name") == quality:
sources = stream.get("sources", [])
non_yospace_sources = [
s for s in sources if not s.get("isYospace", False)
]
non_yospace_sources = [s for s in sources if not s.get("isYospace", False)]
if non_yospace_sources:
selected_url = non_yospace_sources[0].get("url")
@@ -361,9 +353,7 @@ class RTLPlusProvider(StreamingProvider):
if not license_url:
continue
def create_drm_config(
drm_system, priority, server_url, headers
):
def create_drm_config(drm_system, priority, server_url, headers):
return DRMConfig(
system=drm_system,
priority=priority,
@@ -407,14 +397,10 @@ class RTLPlusProvider(StreamingProvider):
return drm_configs
except requests.RequestException as e:
logger.error(
f"Error fetching DRM configs for RTL+ channel {channel_id}: {e}"
)
logger.error(f"Error fetching DRM configs for RTL+ channel {channel_id}: {e}")
return []
except Exception as e:
logger.error(
f"Error parsing DRM configs for RTL+ channel {channel_id}: {e}"
)
logger.error(f"Error parsing DRM configs for RTL+ channel {channel_id}: {e}")
return []
@staticmethod
+279 -33
View File
@@ -2,6 +2,7 @@
import os
import sys
import threading
import traceback
from datetime import datetime
import time
import json
@@ -1031,60 +1032,305 @@ class UltimateService:
@self.app.route('/api/providers/<provider>/channels/<channel_id>/pssh')
def get_channel_pssh(provider, channel_id):
"""
Extract PSSH data for a channel.
Query parameters:
- country: Optional country code for geo-specific manifests
- force_refresh: If 'true', bypass cache and re-extract PSSH
Returns:
{
"provider": "provider_name",
"channel_id": "channel_id",
"manifest_url": "https://...",
"pssh_data": [
{
"system_id": "edef8ba9-79d6-4ace-a3c8-27dcd51d21ed",
"drm_system": "com.widevine.alpha",
"pssh_box": "AAAANHBzc2g...",
"key_ids": ["64656d6f..."],
"source": "mp4_segment"
}
],
"count": 1,
"cached": true
}
"""
try:
# Get the manifest URL first
manifest_url = self.manager.get_channel_manifest(
provider_name=provider,
channel_id=channel_id,
country=request.query.get('country')
)
# Parse query parameters
country = request.params.get('country')
force_refresh = request.params.get('force_refresh', '').lower() == 'true'
# Clear cache if force refresh requested
if force_refresh:
cache_key = f"{provider}:{channel_id}"
self.manager.drm_operations.pssh_cache.clear()
logger.info(f"Cache cleared for force_refresh request: {cache_key}")
# Get manifest URL
try:
manifest_url = self.manager.get_channel_manifest(
provider_name=provider,
channel_id=channel_id,
country=country
)
except ValueError as e:
response.status = 404
return {
'error': 'Provider not found',
'message': str(e),
'provider': provider
}
except Exception as e:
response.status = 500
return {
'error': 'Failed to get manifest',
'message': str(e),
'provider': provider,
'channel_id': channel_id
}
if not manifest_url:
response.status = 404
return {'error': f'Manifest not available for channel "{channel_id}" from provider "{provider}"'}
return {
'error': 'Manifest not available',
'message': f'No manifest found for channel "{channel_id}" from provider "{provider}"',
'provider': provider,
'channel_id': channel_id
}
# Extract PSSH data from the manifest
pssh_data_list = self.manager.extract_pssh_from_manifest(manifest_url)
# Check if we're using cache
cache_key = f"{provider}:{channel_id}"
cached_pssh = self.manager.drm_operations.pssh_cache.get(cache_key)
was_cached = cached_pssh is not None
# Extract PSSH data (uses cache internally unless force_refresh)
try:
pssh_data_list = self.manager.drm_operations._extract_pssh_from_manifest(
manifest_url
)
except Exception as e:
response.status = 500
return {
'error': 'Failed to extract PSSH',
'message': str(e),
'provider': provider,
'channel_id': channel_id,
'manifest_url': manifest_url,
'traceback': traceback.format_exc() if self.app.config.get('debug') else None
}
if not pssh_data_list:
response.status = 404
return {
'error': f'No PSSH data found in manifest for channel "{channel_id}" from provider "{provider}"'}
'error': 'No PSSH data found',
'message': f'No PSSH data found in manifest or segments for channel "{channel_id}"',
'provider': provider,
'channel_id': channel_id,
'manifest_url': manifest_url,
'hint': 'This channel may not use DRM, or PSSH extraction failed'
}
# Convert PSSH data to dictionary format for JSON response
# Convert PSSH data to dictionary format
pssh_list = []
for pssh_data in pssh_data_list:
if hasattr(pssh_data, 'to_dict'):
pssh_list.append(pssh_data.to_dict())
else:
# Fallback for basic PSSH data structure
pssh_dict = {
'pssh': getattr(pssh_data, 'pssh', str(pssh_data)) if hasattr(pssh_data, 'pssh') else str(
pssh_data),
'system_id': getattr(pssh_data, 'system_id', None),
'key_id': getattr(pssh_data, 'key_id', None) if hasattr(pssh_data, 'key_id') else None
}
# Remove None values
pssh_dict = {k: v for k, v in pssh_dict.items() if v is not None}
pssh_list.append(pssh_dict)
pssh_dict = {
'system_id': pssh_data.system_id,
'drm_system': pssh_data.drm_system.value if pssh_data.drm_system else None,
'pssh_box': pssh_data.pssh_box if pssh_data.pssh_box else None,
'key_ids': pssh_data.key_ids if pssh_data.key_ids else [],
'source': pssh_data.source
}
# Add human-readable system name
if pssh_data.drm_system:
pssh_dict['drm_system_name'] = {
'com.widevine.alpha': 'Widevine',
'com.microsoft.playready': 'PlayReady',
'com.apple.fps': 'FairPlay',
'org.w3.clearkey': 'ClearKey',
'com.huawei.wiseplay': 'Wiseplay'
}.get(pssh_data.drm_system.value, pssh_data.drm_system.value)
# Remove None values for cleaner response
pssh_dict = {k: v for k, v in pssh_dict.items() if v is not None}
pssh_list.append(pssh_dict)
response.status = 200
return {
'provider': provider,
'channel_id': channel_id,
'manifest_url': manifest_url,
'pssh_data': pssh_list,
'count': len(pssh_list)
'count': len(pssh_list),
'cached': was_cached,
'cache_ttl_seconds': self.manager.drm_operations.pssh_cache.ttl
}
except ValueError as val_err:
# This handles the case where manager raises ValueError for unknown provider
logger.error(f"API Error in /api/providers/{provider}/channels/{channel_id}/pssh: {str(val_err)}")
response.status = 404
return {'error': str(val_err)}
except Exception as api_err:
logger.error(f"API Error in /api/providers/{provider}/channels/{channel_id}/pssh: {str(api_err)}")
except Exception as e:
# Catch-all for unexpected errors
logger.error(f"Unexpected error in get_channel_pssh: {e}")
logger.error(traceback.format_exc())
response.status = 500
return {'error': f'Internal server error: {str(api_err)}'}
return {
'error': 'Internal server error',
'message': str(e),
'provider': provider,
'channel_id': channel_id,
'traceback': traceback.format_exc() if self.app.config.get('debug') else None
}
@self.app.route('/api/providers/<provider>/channels/<channel_id>/pssh/refresh', method='POST')
def refresh_channel_pssh(provider, channel_id):
"""
Force refresh PSSH data for a channel (clears cache and re-extracts).
This is useful when:
- Keys have been rotated
- Manifest structure has changed
- Previous extraction failed
"""
try:
# Clear cache for this specific channel
cache_key = f"{provider}:{channel_id}"
# Check if entry exists in cache
cached = self.manager.drm_operations.pssh_cache.get(cache_key)
# Clear it
if cached:
# Remove specific key (you may need to add this method to PSSHCache)
with self.manager.drm_operations.pssh_cache.lock:
if cache_key in self.manager.drm_operations.pssh_cache.cache:
del self.manager.drm_operations.pssh_cache.cache[cache_key]
logger.info(f"Cleared cache for {cache_key}")
# Now extract fresh data
country = request.params.get('country')
manifest_url = self.manager.get_channel_manifest(
provider_name=provider,
channel_id=channel_id,
country=country
)
if not manifest_url:
response.status = 404
return {
'error': 'Manifest not available',
'provider': provider,
'channel_id': channel_id
}
# Extract and cache
pssh_data_list = self.manager.drm_operations._extract_pssh_from_manifest(
manifest_url
)
if pssh_data_list:
# Cache it
self.manager.drm_operations.pssh_cache.set(cache_key, pssh_data_list)
response.status = 200
return {
'message': 'PSSH data refreshed',
'provider': provider,
'channel_id': channel_id,
'count': len(pssh_data_list),
'was_cached': cached is not None,
'now_cached': len(pssh_data_list) > 0
}
except Exception as e:
logger.error(f"Error refreshing PSSH: {e}")
response.status = 500
return {
'error': 'Failed to refresh',
'message': str(e)
}
@self.app.route('/api/cache/pssh', method='DELETE')
def clear_pssh_cache():
"""
Clear all PSSH cache entries.
This is useful for:
- Debugging
- Freeing memory
- Forcing re-extraction of all channels
"""
try:
cache_size = len(self.manager.drm_operations.pssh_cache.cache)
self.manager.drm_operations.pssh_cache.clear()
response.status = 200
return {
'message': 'PSSH cache cleared',
'entries_cleared': cache_size
}
except Exception as e:
logger.error(f"Error clearing cache: {e}")
response.status = 500
return {
'error': 'Failed to clear cache',
'message': str(e)
}
@self.app.route('/api/cache/pssh', method='GET')
def get_pssh_cache_stats():
"""
Get PSSH cache statistics.
Returns information about:
- Number of cached entries
- TTL configuration
- Memory usage estimate
"""
try:
cache = self.manager.drm_operations.pssh_cache
with cache.lock:
entries = []
total_size = 0
for key, (pssh_list, timestamp) in cache.cache.items():
import time
age = time.time() - timestamp
expires_in = cache.ttl - age
# Estimate size
size = sum(
len(p.pssh_box) + len(str(p.key_ids)) + len(p.system_id)
for p in pssh_list
)
total_size += size
entries.append({
'key': key,
'pssh_count': len(pssh_list),
'age_seconds': int(age),
'expires_in_seconds': int(expires_in),
'size_bytes': size
})
response.status = 200
return {
'total_entries': len(entries),
'ttl_seconds': cache.ttl,
'total_size_bytes': total_size,
'total_size_mb': round(total_size / 1024 / 1024, 2),
'entries': sorted(entries, key=lambda x: x['age_seconds'], reverse=True)
}
except Exception as e:
logger.error(f"Error getting cache stats: {e}")
response.status = 500
return {
'error': 'Failed to get cache stats',
'message': str(e)
}
@self.app.route('/api/providers/<provider>/channels/<channel_id>/epg')
def get_channel_epg(provider, channel_id):