Skip to content

Commit 33e347a

Browse files
Fix embed params being dropped in page swaps (#7918)
1 parent 899cbbd commit 33e347a

10 files changed

Lines changed: 296 additions & 22 deletions

File tree

‎e2e_playwright/multipage_apps/mpa_basics_test.py‎

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -176,3 +176,20 @@ def test_removes_query_params_when_swapping_pages(page: Page, app_port: int):
176176
wait_for_app_run(page)
177177

178178
assert page.url == f"http://localhost:{app_port}/page3"
179+
180+
181+
def test_removes_non_embed_query_params_when_swapping_pages(page: Page, app_port: int):
182+
"""Test that query params are removed when swapping pages"""
183+
184+
page.goto(
185+
f"http://localhost:{app_port}/page_7?foo=bar&embed=True&embed_options=show_toolbar&embed_options=show_colored_line"
186+
)
187+
wait_for_app_loaded(page)
188+
189+
page.get_by_test_id("stSidebarNav").locator("a").nth(2).click()
190+
wait_for_app_run(page)
191+
192+
assert (
193+
page.url
194+
== f"http://localhost:{app_port}/page3?embed=true&embed_options=show_toolbar&embed_options=show_colored_line"
195+
)

‎frontend/app/src/App.test.tsx‎

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1193,6 +1193,48 @@ describe("App.sendRerunBackMsg", () => {
11931193
queryParams: "",
11941194
})
11951195
})
1196+
1197+
it("retains embed query params even if the page hash is different", () => {
1198+
const embedParams =
1199+
"embed=true&embed_options=disable_scrolling&embed_options=show_colored_line"
1200+
1201+
const prevWindowLocation = window.location
1202+
// @ts-expect-error
1203+
delete window.location
1204+
// @ts-expect-error
1205+
window.location = {
1206+
assign: jest.fn(),
1207+
search: `foo=bar&${embedParams}`,
1208+
}
1209+
1210+
wrapper.setState({
1211+
currentPageScriptHash: "current_page_hash",
1212+
queryParams: `foo=bar&${embedParams}`,
1213+
})
1214+
const sendMessageFunc = jest.spyOn(
1215+
// @ts-expect-error
1216+
instance.hostCommunicationMgr,
1217+
"sendMessageToHost"
1218+
)
1219+
1220+
instance.sendRerunBackMsg(undefined, "some_other_page_hash")
1221+
1222+
// @ts-expect-error
1223+
expect(instance.sendBackMsg).toHaveBeenCalledWith({
1224+
rerunScript: {
1225+
pageScriptHash: "some_other_page_hash",
1226+
pageName: "",
1227+
queryString: embedParams,
1228+
},
1229+
})
1230+
1231+
expect(sendMessageFunc).toHaveBeenCalledWith({
1232+
type: "SET_QUERY_PARAM",
1233+
queryParams: embedParams,
1234+
})
1235+
1236+
window.location = prevWindowLocation
1237+
})
11961238
})
11971239

11981240
// * handlePageNotFound has branching error messages depending on pageName

‎frontend/app/src/App.tsx‎

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -124,6 +124,7 @@ import withScreencast, {
124124

125125
// Used to import fonts + responsive reboot items
126126
import "@streamlit/app/src/assets/css/theme.scss"
127+
import { preserveEmbedQueryParams } from "@streamlit/lib/src/util/utils"
127128

128129
export interface Props {
129130
screenCast: ScreenCastHOC
@@ -865,10 +866,14 @@ export class App extends PureComponent<Props, State> {
865866
// e.g. the case where the user clicks the back button.
866867
// See https://github.com/streamlit/streamlit/pull/6271#issuecomment-1465090690 for the discussion.
867868
if (prevPageName !== newPageName) {
869+
// If embed params need to be changed, make sure to change to other parts of the code that reference preserveEmbedQueryParams
870+
const queryString = preserveEmbedQueryParams()
871+
const qs = queryString ? `?${queryString}` : ""
872+
868873
const basePathPrefix = basePath ? `/${basePath}` : ""
869874

870875
const pagePath = viewingMainPage ? "" : newPageName
871-
const pageUrl = `${basePathPrefix}/${pagePath}`
876+
const pageUrl = `${basePathPrefix}/${pagePath}${qs}`
872877

873878
window.history.pushState({}, "", pageUrl)
874879
}
@@ -1307,8 +1312,8 @@ export class App extends PureComponent<Props, State> {
13071312
// The user specified exactly which page to run. We can simply use this
13081313
// value in the BackMsg we send to the server.
13091314
if (pageScriptHash != currentPageScriptHash) {
1310-
// clear query parameters within a page change
1311-
queryString = ""
1315+
// clear non-embed query parameters within a page change
1316+
queryString = preserveEmbedQueryParams()
13121317
this.hostCommunicationMgr.sendMessageToHost({
13131318
type: "SET_QUERY_PARAM",
13141319
queryParams: queryString,

‎frontend/lib/src/util/utils.test.ts‎

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@ import {
2323
getLoadingScreenType,
2424
isEmbed,
2525
setCookie,
26+
preserveEmbedQueryParams,
2627
} from "./utils"
2728

2829
describe("getCookie", () => {
@@ -331,4 +332,47 @@ describe("getLoadingScreenType", () => {
331332

332333
expect(getLoadingScreenType()).toBe(LoadingScreenType.V2)
333334
})
335+
336+
describe("preserveEmbedQueryParams", () => {
337+
let prevWindowLocation: Location
338+
afterEach(() => {
339+
window.location = prevWindowLocation
340+
})
341+
342+
it("should return an empty string if not in embed mode", () => {
343+
// @ts-expect-error
344+
delete window.location
345+
// @ts-expect-error
346+
window.location = {
347+
assign: jest.fn(),
348+
search: "foo=bar",
349+
}
350+
expect(preserveEmbedQueryParams()).toBe("")
351+
})
352+
353+
it("should preserve embed query string even with no embed options and remove foo=bar", () => {
354+
// @ts-expect-error
355+
delete window.location
356+
// @ts-expect-error
357+
window.location = {
358+
assign: jest.fn(),
359+
search: "embed=true&foo=bar",
360+
}
361+
expect(preserveEmbedQueryParams()).toBe("embed=true")
362+
})
363+
364+
it("should preserve embed query string with embed options and remove foo=bar", () => {
365+
// @ts-expect-error
366+
delete window.location
367+
// @ts-expect-error
368+
window.location = {
369+
assign: jest.fn(),
370+
search:
371+
"embed=true&embed_options=option1&embed_options=option2&foo=bar",
372+
}
373+
expect(preserveEmbedQueryParams()).toBe(
374+
"embed=true&embed_options=option1&embed_options=option2"
375+
)
376+
})
377+
})
334378
})

‎frontend/lib/src/util/utils.ts‎

Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -98,6 +98,31 @@ export function getEmbedUrlParams(embedKey: string): Set<string> {
9898
return embedUrlParams
9999
}
100100

101+
/**
102+
* Returns "embed" and "embed_options" query param options in the url. Returns empty string if not embedded.
103+
* Example:
104+
* returns "embed=true&embed_options=show_loading_screen_v2" if the url is
105+
* http://localhost:3000/test?embed=true&embed_options=show_loading_screen_v2
106+
*/
107+
export function preserveEmbedQueryParams(): string {
108+
if (!isEmbed()) {
109+
return ""
110+
}
111+
112+
const embedOptionsValues = new URLSearchParams(
113+
window.location.search
114+
).getAll(EMBED_OPTIONS_QUERY_PARAM_KEY)
115+
116+
// instantiate multiple key values with an array of string pairs
117+
// https://stackoverflow.com/questions/72571132/urlsearchparams-with-multiple-values
118+
const embedUrlMap: string[][] = []
119+
embedUrlMap.push([EMBED_QUERY_PARAM_KEY, EMBED_TRUE])
120+
embedOptionsValues.forEach((embedValue: string) => {
121+
embedUrlMap.push([EMBED_OPTIONS_QUERY_PARAM_KEY, embedValue])
122+
})
123+
return new URLSearchParams(embedUrlMap).toString()
124+
}
125+
101126
/**
102127
* Returns true if the URL parameters contain ?embed=true (case insensitive).
103128
*/

‎lib/streamlit/commands/execution_control.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -157,7 +157,7 @@ def switch_page(page: str) -> NoReturn: # type: ignore[misc]
157157

158158
ctx.script_requests.request_rerun(
159159
RerunData(
160-
query_string="",
160+
query_string=ctx.query_string,
161161
page_script_hash=matched_pages[0]["page_script_hash"],
162162
)
163163
)

‎lib/streamlit/commands/experimental_query_params.py‎

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -16,15 +16,16 @@
1616
from typing import Any, Dict, List, Union
1717

1818
from streamlit import util
19+
from streamlit.constants import (
20+
EMBED_OPTIONS_QUERY_PARAM,
21+
EMBED_QUERY_PARAM,
22+
EMBED_QUERY_PARAMS_KEYS,
23+
)
1924
from streamlit.errors import StreamlitAPIException
2025
from streamlit.proto.ForwardMsg_pb2 import ForwardMsg
2126
from streamlit.runtime.metrics_util import gather_metrics
2227
from streamlit.runtime.scriptrunner import get_script_run_ctx
2328

24-
EMBED_QUERY_PARAM = "embed"
25-
EMBED_OPTIONS_QUERY_PARAM = "embed_options"
26-
EMBED_QUERY_PARAMS_KEYS = [EMBED_QUERY_PARAM, EMBED_OPTIONS_QUERY_PARAM]
27-
2829

2930
@gather_metrics("experimental_get_query_params")
3031
def get_query_params() -> Dict[str, List[str]]:

‎lib/streamlit/constants.py‎

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,17 @@
1+
# Copyright (c) Streamlit Inc. (2018-2022) Snowflake Inc. (2022-2024)
2+
#
3+
# Licensed under the Apache License, Version 2.0 (the "License");
4+
# you may not use this file except in compliance with the License.
5+
# You may obtain a copy of the License at
6+
#
7+
# http://www.apache.org/licenses/LICENSE-2.0
8+
#
9+
# Unless required by applicable law or agreed to in writing, software
10+
# distributed under the License is distributed on an "AS IS" BASIS,
11+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+
# See the License for the specific language governing permissions and
13+
# limitations under the License.
14+
15+
EMBED_QUERY_PARAM = "embed"
16+
EMBED_OPTIONS_QUERY_PARAM = "embed_options"
17+
EMBED_QUERY_PARAMS_KEYS = [EMBED_QUERY_PARAM, EMBED_OPTIONS_QUERY_PARAM]

‎lib/streamlit/runtime/state/query_params.py‎

Lines changed: 34 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,9 @@
1414

1515
from dataclasses import dataclass, field
1616
from typing import Dict, Iterable, Iterator, List, MutableMapping, Union
17+
from urllib import parse
1718

19+
from streamlit.constants import EMBED_QUERY_PARAMS_KEYS
1820
from streamlit.errors import StreamlitAPIException
1921
from streamlit.proto.ForwardMsg_pb2 import ForwardMsg
2022

@@ -29,7 +31,12 @@ class QueryParams(MutableMapping[str, str]):
2931

3032
def __iter__(self) -> Iterator[str]:
3133
self._ensure_single_query_api_used()
32-
return iter(self._query_params.keys())
34+
35+
return iter(
36+
key
37+
for key in self._query_params.keys()
38+
if key not in EMBED_QUERY_PARAMS_KEYS
39+
)
3340

3441
def __getitem__(self, key: str) -> str:
3542
"""Retrieves a value for a given key in query parameters.
@@ -38,6 +45,8 @@ def __getitem__(self, key: str) -> str:
3845
"""
3946
self._ensure_single_query_api_used()
4047
try:
48+
if key in EMBED_QUERY_PARAMS_KEYS:
49+
raise KeyError(missing_key_error_message(key))
4150
value = self._query_params[key]
4251
if isinstance(value, list):
4352
if len(value) == 0:
@@ -56,6 +65,10 @@ def __setitem__(self, key: str, value: Union[str, Iterable[str]]) -> None:
5665
f"You cannot set a query params key `{key}` to a dictionary."
5766
)
5867

68+
if key in EMBED_QUERY_PARAMS_KEYS:
69+
raise StreamlitAPIException(
70+
"Query param embed and embed_options (case-insensitive) cannot be set programmatically."
71+
)
5972
# Type checking users should handle the string serialization themselves
6073
# We will accept any type for the list and serialize to str just in case
6174
if isinstance(value, Iterable) and not isinstance(value, str):
@@ -66,25 +79,28 @@ def __setitem__(self, key: str, value: Union[str, Iterable[str]]) -> None:
6679

6780
def __delitem__(self, key: str) -> None:
6881
try:
82+
if key in EMBED_QUERY_PARAMS_KEYS:
83+
raise KeyError(missing_key_error_message(key))
6984
del self._query_params[key]
7085
self._send_query_param_msg()
7186
except KeyError:
7287
raise KeyError(missing_key_error_message(key))
7388

7489
def get_all(self, key: str) -> List[str]:
7590
self._ensure_single_query_api_used()
76-
if key not in self._query_params:
91+
if key not in self._query_params or key in EMBED_QUERY_PARAMS_KEYS:
7792
return []
7893
value = self._query_params[key]
7994
return value if isinstance(value, list) else [value]
8095

8196
def __len__(self) -> int:
8297
self._ensure_single_query_api_used()
83-
return len(self._query_params)
98+
return len(
99+
{key for key in self._query_params if key not in EMBED_QUERY_PARAMS_KEYS}
100+
)
84101

85102
def _send_query_param_msg(self) -> None:
86103
# Avoid circular imports
87-
from streamlit.commands.experimental_query_params import _ensure_no_embed_params
88104
from streamlit.runtime.scriptrunner import get_script_run_ctx
89105

90106
ctx = get_script_run_ctx()
@@ -93,27 +109,31 @@ def _send_query_param_msg(self) -> None:
93109
self._ensure_single_query_api_used()
94110

95111
msg = ForwardMsg()
96-
msg.page_info_changed.query_string = _ensure_no_embed_params(
97-
self._query_params, ctx.query_string
112+
msg.page_info_changed.query_string = parse.urlencode(
113+
self._query_params, doseq=True
98114
)
99115
ctx.query_string = msg.page_info_changed.query_string
100116
ctx.enqueue(msg)
101117

102118
def clear(self) -> None:
103-
self._query_params.clear()
119+
new_query_params = {}
120+
for key, value in self._query_params.items():
121+
if key in EMBED_QUERY_PARAMS_KEYS:
122+
new_query_params[key] = value
123+
self._query_params = new_query_params
124+
104125
self._send_query_param_msg()
105126

106127
def to_dict(self) -> Dict[str, str]:
107128
self._ensure_single_query_api_used()
108-
# return the last query param if multiple keys are set
109-
return {key: self[key] for key in self._query_params}
129+
# return the last query param if multiple values are set
130+
return {
131+
key: self[key]
132+
for key in self._query_params
133+
if key not in EMBED_QUERY_PARAMS_KEYS
134+
}
110135

111136
def set_with_no_forward_msg(self, key: str, val: Union[List[str], str]) -> None:
112-
# Avoid circular imports
113-
from streamlit.commands.experimental_query_params import EMBED_QUERY_PARAMS_KEYS
114-
115-
if key.lower() in EMBED_QUERY_PARAMS_KEYS:
116-
return
117137
self._query_params[key] = val
118138

119139
def clear_with_no_forward_msg(self) -> None:

0 commit comments

Comments
 (0)