Skip to content

Commit 90e84a0

Browse files
committed
fix(auth): address mypy guards and pyink in PKCE flow
1 parent 8614cb4 commit 90e84a0

2 files changed

Lines changed: 22 additions & 21 deletions

File tree

src/google/adk/auth/auth_handler.py

Lines changed: 16 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -70,7 +70,7 @@ async def parse_and_store_auth_response(self, state: State) -> None:
7070
state[credential_key] = await self.exchange_auth_token()
7171

7272
def _validate(self) -> None:
73-
if not self.auth_scheme:
73+
if not self.auth_config.auth_scheme:
7474
raise ValueError("auth_scheme is empty.")
7575

7676
def get_auth_response(self, state: State) -> AuthCredential:
@@ -160,7 +160,8 @@ def generate_auth_uri(
160160
auth_scheme = self.auth_config.auth_scheme
161161
auth_credential = self.auth_config.raw_auth_credential
162162
if not auth_credential or not auth_credential.oauth2:
163-
raise ValueError("raw_auth_credential or oauth2 is empty")
163+
raise ValueError("OAuth2 auth_credential with oauth2 config is required.")
164+
oauth2_credential = auth_credential.oauth2
164165

165166
if isinstance(auth_scheme, OpenIdConnectWithConfig):
166167
authorization_endpoint = auth_scheme.authorization_endpoint
@@ -189,24 +190,24 @@ def generate_auth_uri(
189190
scopes = list(scopes.keys())
190191

191192
client = OAuth2Session(
192-
auth_credential.oauth2.client_id,
193-
auth_credential.oauth2.client_secret,
193+
oauth2_credential.client_id,
194+
oauth2_credential.client_secret,
194195
scope=" ".join(scopes),
195-
redirect_uri=auth_credential.oauth2.redirect_uri,
196-
code_challenge_method=auth_credential.oauth2.code_challenge_method,
196+
redirect_uri=oauth2_credential.redirect_uri,
197+
code_challenge_method=oauth2_credential.code_challenge_method,
197198
)
198199
params = {
199200
"access_type": "offline",
200201
"prompt": "consent",
201202
}
202-
if auth_credential.oauth2.audience:
203-
params["audience"] = auth_credential.oauth2.audience
203+
if oauth2_credential.audience:
204+
params["audience"] = oauth2_credential.audience
204205

205206
# If using PKCE with S256, ensure a code_verifier exists.
206207
# If not provided in the credential, generate a cryptographically secure
207208
# random token of 48 characters (OAuth2 recommends 43-128 characters).
208-
code_verifier = auth_credential.oauth2.code_verifier
209-
method = auth_credential.oauth2.code_challenge_method
209+
code_verifier = oauth2_credential.code_verifier
210+
method = oauth2_credential.code_challenge_method
210211

211212
if method:
212213
if method != "S256":
@@ -222,9 +223,10 @@ def generate_auth_uri(
222223
)
223224

224225
exchanged_auth_credential = auth_credential.model_copy(deep=True)
225-
exchanged_auth_credential.oauth2.auth_uri = uri
226-
exchanged_auth_credential.oauth2.state = state
227-
if code_verifier:
228-
exchanged_auth_credential.oauth2.code_verifier = code_verifier
226+
if exchanged_auth_credential.oauth2:
227+
exchanged_auth_credential.oauth2.auth_uri = uri
228+
exchanged_auth_credential.oauth2.state = state
229+
if code_verifier:
230+
exchanged_auth_credential.oauth2.code_verifier = code_verifier
229231

230232
return exchanged_auth_credential

tests/unittests/auth/test_auth_handler.py

Lines changed: 6 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -66,13 +66,12 @@ def create_authorization_url(self, url, **kwargs):
6666
params = f"client_id={self.client_id}&scope={self.scope}"
6767
if kwargs.get("audience"):
6868
params += f"&audience={kwargs.get('audience')}"
69-
if kwargs.get("code_challenge_method"):
70-
params += (
71-
"&code_challenge_method="
72-
f"{kwargs.get('code_challenge_method')}"
73-
)
74-
if kwargs.get("code_challenge"):
75-
params += f"&code_challenge={kwargs.get('code_challenge')}"
69+
code_challenge_method = self.extra_kwargs.get(
70+
"code_challenge_method"
71+
) or kwargs.get("code_challenge_method")
72+
if code_challenge_method:
73+
params += f"&code_challenge_method={code_challenge_method}"
74+
params += "&code_challenge=mock_code_challenge"
7675
return f"{url}?{params}", "mock_state"
7776

7877
def fetch_token(

0 commit comments

Comments
 (0)