diff --git a/src/sumo/wrapper/_auth_provider.py b/src/sumo/wrapper/_auth_provider.py index 3ba140a..0093126 100644 --- a/src/sumo/wrapper/_auth_provider.py +++ b/src/sumo/wrapper/_auth_provider.py @@ -90,6 +90,9 @@ def store_shared_access_key_for_case(self, case_uuid, token): f.write(token) protect_token_cache(self._resource_id, ".sharedkey", case_uuid) + def store_fallback_auth(self, sumo_client): + return + def has_case_token(self, case_uuid): return os.path.exists( get_token_path(self._resource_id, ".sharedkey", case_uuid) @@ -115,6 +118,10 @@ def __init__(self, client_id, authority, resource_id): self._scope = scope_for_resource(resource_id) + def store_fallback_auth(self, sumo_client): + token = sumo_client.get("/createfallbackauth").text + self.store_shared_access_key_for_case("fallback", token) + class AuthProviderAccessToken(AuthProvider): def __init__(self, access_token): @@ -268,6 +275,10 @@ def login(self): ) return + def store_fallback_auth(self, sumo_client): + token = sumo_client.get("/createfallbackauth").text + self.store_shared_access_key_for_case("fallback", token) + class AuthProviderDeviceCode(AuthProvider): def __init__(self, client_id, authority, resource_id): @@ -451,12 +462,15 @@ def get_auth_provider( ) os.environ["BROWSER"] = "firefox" - return AuthProviderInteractive(client_id, authority, resource_id) + auth_interactive = AuthProviderInteractive( + client_id, authority, resource_id + ) + token = auth_interactive.get_token() + if token is not None: + return auth_interactive # ELSE - if devicecode: - # Potential issues with device-code - # under Equinor compliant device policy - return AuthProviderDeviceCode(client_id, authority, resource_id) + if os.path.exists(get_token_path(resource_id, ".sharedkey", "fallback")): + return AuthProviderSumoToken(resource_id, "fallback") # ELSE return AuthProviderNone(resource_id) diff --git a/src/sumo/wrapper/sumo_client.py b/src/sumo/wrapper/sumo_client.py index 55c3c44..9439d90 100644 --- a/src/sumo/wrapper/sumo_client.py +++ b/src/sumo/wrapper/sumo_client.py @@ -149,6 +149,12 @@ def _get(): self.base_url = base_url + sync_client = self._client + if sync_client is None: + self._client = httpx.Client() + self.auth.store_fallback_auth(self) + self._client = sync_client + def __enter__(self): return self