Skip to content

Commit 995cd1c

Browse files
GWealecopybara-github
authored andcommitted
fix: add protection for arbitrary module imports
Close#4947 Co-authored-by: George Weale <gweale@google.com> PiperOrigin-RevId: 888296476
1 parent 0f351bf commit 995cd1c

3 files changed

Lines changed: 292 additions & 3 deletions

File tree

‎src/google/adk/cli/adk_web_server.py‎

Lines changed: 178 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@
2020
importjson
2121
importlogging
2222
importos
23+
importre
2324
importsys
2425
importtime
2526
importtraceback
@@ -138,6 +139,158 @@ def _parse_cors_origins(
138139
returnliteral_origins, combined_regex
139140

140141

142+
def_is_origin_allowed(
143+
origin: str,
144+
allowed_literal_origins: list[str],
145+
allowed_origin_regex: Optional[re.Pattern[str]],
146+
) ->bool:
147+
"""Check whether the given origin matches the allowed origins."""
148+
if"*"inallowed_literal_origins:
149+
returnTrue
150+
iforigininallowed_literal_origins:
151+
returnTrue
152+
ifallowed_origin_regexisnotNone:
153+
returnallowed_origin_regex.fullmatch(origin) isnotNone
154+
returnFalse
155+
156+
157+
def_normalize_origin_scheme(scheme: str) ->str:
158+
"""Normalize request schemes to the browser Origin scheme space."""
159+
ifscheme=="ws":
160+
return"http"
161+
ifscheme=="wss":
162+
return"https"
163+
returnscheme
164+
165+
166+
def_strip_optional_quotes(value: str) ->str:
167+
"""Strip a single pair of wrapping quotes from a header value."""
168+
iflen(value) >=2andvalue[0] =='"'andvalue[-1] =='"':
169+
returnvalue[1:-1]
170+
returnvalue
171+
172+
173+
def_get_scope_header(
174+
scope: dict[str, Any], header_name: bytes
175+
) ->Optional[str]:
176+
"""Return the first matching header value from an ASGI scope."""
177+
forcandidate_name, candidate_valueinscope.get("headers", []):
178+
ifcandidate_name==header_name:
179+
returncandidate_value.decode("latin-1").split(",", 1)[0].strip()
180+
returnNone
181+
182+
183+
def_get_request_origin(scope: dict[str, Any]) ->Optional[str]:
184+
"""Compute the effective origin for the current HTTP/WebSocket request."""
185+
forwarded=_get_scope_header(scope, b"forwarded")
186+
ifforwardedisnotNone:
187+
proto=None
188+
host=None
189+
forelementinforwarded.split(",", 1)[0].split(";"):
190+
if"="notinelement:
191+
continue
192+
name, value=element.split("=", 1)
193+
ifname.strip().lower() =="proto":
194+
proto=_strip_optional_quotes(value.strip())
195+
elifname.strip().lower() =="host":
196+
host=_strip_optional_quotes(value.strip())
197+
ifprotoisnotNoneandhostisnotNone:
198+
returnf"{_normalize_origin_scheme(proto)}://{host}"
199+
200+
host=_get_scope_header(scope, b"x-forwarded-host")
201+
ifhostisNone:
202+
host=_get_scope_header(scope, b"host")
203+
ifhostisNone:
204+
returnNone
205+
206+
proto=_get_scope_header(scope, b"x-forwarded-proto")
207+
ifprotoisNone:
208+
proto=scope.get("scheme", "http")
209+
returnf"{_normalize_origin_scheme(proto)}://{host}"
210+
211+
212+
def_is_request_origin_allowed(
213+
origin: str,
214+
scope: dict[str, Any],
215+
allowed_literal_origins: list[str],
216+
allowed_origin_regex: Optional[re.Pattern[str]],
217+
has_configured_allowed_origins: bool,
218+
) ->bool:
219+
"""Validate an Origin header against explicit config or same-origin."""
220+
ifhas_configured_allowed_originsand_is_origin_allowed(
221+
origin, allowed_literal_origins, allowed_origin_regex
222+
):
223+
returnTrue
224+
225+
request_origin=_get_request_origin(scope)
226+
ifrequest_originisNone:
227+
returnFalse
228+
returnorigin==request_origin
229+
230+
231+
_SAFE_HTTP_METHODS=frozenset({"GET", "HEAD", "OPTIONS"})
232+
233+
234+
class_OriginCheckMiddleware:
235+
"""ASGI middleware that blocks cross-origin state-changing requests."""
236+
237+
def__init__(
238+
self,
239+
app: Any,
240+
has_configured_allowed_origins: bool,
241+
allowed_origins: list[str],
242+
allowed_origin_regex: Optional[re.Pattern[str]],
243+
) ->None:
244+
self._app=app
245+
self._has_configured_allowed_origins=has_configured_allowed_origins
246+
self._allowed_origins=allowed_origins
247+
self._allowed_origin_regex=allowed_origin_regex
248+
249+
asyncdef__call__(
250+
self,
251+
scope: dict[str, Any],
252+
receive: Any,
253+
send: Any,
254+
) ->None:
255+
ifscope["type"] !="http":
256+
awaitself._app(scope, receive, send)
257+
return
258+
259+
method=scope.get("method", "GET")
260+
ifmethodin_SAFE_HTTP_METHODS:
261+
awaitself._app(scope, receive, send)
262+
return
263+
264+
origin=_get_scope_header(scope, b"origin")
265+
iforiginisNone:
266+
awaitself._app(scope, receive, send)
267+
return
268+
269+
if_is_request_origin_allowed(
270+
origin,
271+
scope,
272+
self._allowed_origins,
273+
self._allowed_origin_regex,
274+
self._has_configured_allowed_origins,
275+
):
276+
awaitself._app(scope, receive, send)
277+
return
278+
279+
response_body=b"Forbidden: origin not allowed"
280+
awaitsend({
281+
"type": "http.response.start",
282+
"status": 403,
283+
"headers": [
284+
(b"content-type", b"text/plain"),
285+
(b"content-length", str(len(response_body)).encode()),
286+
],
287+
})
288+
awaitsend({
289+
"type": "http.response.body",
290+
"body": response_body,
291+
})
292+
293+
141294
classApiServerSpanExporter(export_lib.SpanExporter):
142295

143296
def__init__(self, trace_dict):
@@ -757,8 +910,12 @@ async def internal_lifespan(app: FastAPI):
757910
# Run the FastAPI server.
758911
app=FastAPI(lifespan=internal_lifespan)
759912

913+
has_configured_allowed_origins=bool(allow_origins)
760914
ifallow_origins:
761915
literal_origins, combined_regex=_parse_cors_origins(allow_origins)
916+
compiled_origin_regex= (
917+
re.compile(combined_regex) ifcombined_regexisnotNoneelseNone
918+
)
762919
app.add_middleware(
763920
CORSMiddleware,
764921
allow_origins=literal_origins,
@@ -767,6 +924,16 @@ async def internal_lifespan(app: FastAPI):
767924
allow_methods=["*"],
768925
allow_headers=["*"],
769926
)
927+
else:
928+
literal_origins= []
929+
compiled_origin_regex=None
930+
931+
app.add_middleware(
932+
_OriginCheckMiddleware,
933+
has_configured_allowed_origins=has_configured_allowed_origins,
934+
allowed_origins=literal_origins,
935+
allowed_origin_regex=compiled_origin_regex,
936+
)
770937

771938
@app.get("/health")
772939
asyncdefhealth() ->dict[str, str]:
@@ -1802,14 +1969,23 @@ async def run_agent_live(
18021969
enable_affective_dialog: bool|None=Query(default=None),
18031970
enable_session_resumption: bool|None=Query(default=None),
18041971
) ->None:
1972+
ws_origin=websocket.headers.get("origin")
1973+
ifws_originisnotNoneandnot_is_request_origin_allowed(
1974+
ws_origin,
1975+
websocket.scope,
1976+
literal_origins,
1977+
compiled_origin_regex,
1978+
has_configured_allowed_origins,
1979+
):
1980+
awaitwebsocket.close(code=1008, reason="Origin not allowed")
1981+
return
1982+
18051983
awaitwebsocket.accept()
18061984

18071985
session=awaitself.session_service.get_session(
18081986
app_name=app_name, user_id=user_id, session_id=session_id
18091987
)
18101988
ifnotsession:
1811-
# Accept first so that the client is aware of connection establishment,
1812-
# then close with a specific code.
18131989
awaitwebsocket.close(code=1002, reason="Session not found")
18141990
return
18151991

‎tests/unittests/cli/test_adk_web_server_run_live.py‎

Lines changed: 73 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,7 @@
2222
fromgoogle.adk.events.eventimportEvent
2323
fromgoogle.adk.sessions.in_memory_session_serviceimportInMemorySessionService
2424
importpytest
25+
fromstarlette.websocketsimportWebSocketDisconnect
2526

2627

2728
class_DummyAgent(BaseAgent):
@@ -203,3 +204,75 @@ async def _get_runner_async(_self, _app_name: str):
203204
run_config.session_resumption.transparent
204205
isexpected_session_resumption_transparent
205206
)
207+
208+
209+
_WS_BASE_URL= (
210+
"/run_live"
211+
"?app_name=test_app"
212+
"&user_id=user"
213+
"&session_id=session"
214+
"&modalities=AUDIO"
215+
)
216+
217+
218+
def_build_ws_client():
219+
"""Build a TestClient wired to a capturing runner."""
220+
session_service=InMemorySessionService()
221+
asyncio.run(
222+
session_service.create_session(
223+
app_name="test_app",
224+
user_id="user",
225+
session_id="session",
226+
state={},
227+
)
228+
)
229+
230+
runner=_CapturingRunner()
231+
adk_web_server=AdkWebServer(
232+
agent_loader=_DummyAgentLoader(),
233+
session_service=session_service,
234+
memory_service=types.SimpleNamespace(),
235+
artifact_service=types.SimpleNamespace(),
236+
credential_service=types.SimpleNamespace(),
237+
eval_sets_manager=types.SimpleNamespace(),
238+
eval_set_results_manager=types.SimpleNamespace(),
239+
agents_dir=".",
240+
)
241+
242+
asyncdef_get_runner_async(_self, _app_name: str):
243+
returnrunner
244+
245+
adk_web_server.get_runner_async=_get_runner_async.__get__(adk_web_server) # pytype: disable=attribute-error
246+
247+
fast_api_app=adk_web_server.get_fast_api_app(
248+
setup_observer=lambda_observer, _server: None,
249+
tear_down_observer=lambda_observer, _server: None,
250+
)
251+
returnTestClient(fast_api_app)
252+
253+
254+
deftest_run_live_rejects_disallowed_origin():
255+
client=_build_ws_client()
256+
withpytest.raises(WebSocketDisconnect) asexc_info:
257+
withclient.websocket_connect(
258+
_WS_BASE_URL,
259+
headers={"origin": "https://evil.com"},
260+
) asws:
261+
ws.receive_text()
262+
assertexc_info.value.code==1008
263+
264+
265+
deftest_run_live_allows_matching_origin():
266+
client=_build_ws_client()
267+
withclient.websocket_connect(
268+
_WS_BASE_URL,
269+
headers={"origin": "http://testserver"},
270+
) asws:
271+
_=ws.receive_text()
272+
273+
274+
deftest_run_live_allows_no_origin_header():
275+
"""Non-browser clients (curl, wscat, SDKs) send no Origin header."""
276+
client=_build_ws_client()
277+
withclient.websocket_connect(_WS_BASE_URL) asws:
278+
_=ws.receive_text()

‎tests/unittests/cli/test_fast_api.py‎

Lines changed: 41 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -593,7 +593,7 @@ def builder_test_client(
593593
session_service_uri="",
594594
artifact_service_uri="",
595595
memory_service_uri="",
596-
allow_origins=["*"],
596+
allow_origins=None,
597597
a2a=False,
598598
host="127.0.0.1",
599599
port=8000,
@@ -1595,6 +1595,46 @@ def test_builder_final_save_preserves_tools_and_cleans_tmp(
15951595
assertnottmp_dir.exists() ornotany(tmp_dir.iterdir())
15961596

15971597

1598+
deftest_builder_save_rejects_cross_origin_post(builder_test_client, tmp_path):
1599+
response=builder_test_client.post(
1600+
"/builder/save?tmp=true",
1601+
headers={"origin": "https://evil.com"},
1602+
files=[(
1603+
"files",
1604+
("app/root_agent.yaml", b"name: app\n", "application/x-yaml"),
1605+
)],
1606+
)
1607+
1608+
assertresponse.status_code==403
1609+
assertresponse.text=="Forbidden: origin not allowed"
1610+
assertnot (tmp_path/"app"/"tmp"/"app").exists()
1611+
1612+
1613+
deftest_builder_save_allows_same_origin_post(builder_test_client, tmp_path):
1614+
response=builder_test_client.post(
1615+
"/builder/save?tmp=true",
1616+
headers={"origin": "http://testserver"},
1617+
files=[(
1618+
"files",
1619+
("app/root_agent.yaml", b"name: app\n", "application/x-yaml"),
1620+
)],
1621+
)
1622+
1623+
assertresponse.status_code==200
1624+
assertresponse.json() isTrue
1625+
assert (tmp_path/"app"/"tmp"/"app"/"root_agent.yaml").is_file()
1626+
1627+
1628+
deftest_builder_get_allows_cross_origin_get(builder_test_client):
1629+
response=builder_test_client.get(
1630+
"/builder/app/missing?tmp=true",
1631+
headers={"origin": "https://evil.com"},
1632+
)
1633+
1634+
assertresponse.status_code==200
1635+
assertresponse.text==""
1636+
1637+
15981638
deftest_builder_cancel_deletes_tmp_idempotent(builder_test_client, tmp_path):
15991639
tmp_agent_root=tmp_path/"app"/"tmp"/"app"
16001640
tmp_agent_root.mkdir(parents=True, exist_ok=True)

0 commit comments

Comments
 (0)