-If you are using pygwalker on Jupyter Notebook(version<7) and it can't display properly, please execute code to fix it: `pip install "pygwalker[notebook]" --pre`.(close after 15 seconds)
+If you are using pygwalker on Jupyter Notebook(version<7) and it can't display properly, please execute code to fix it: `pip install "pygwalker[notebook-legacy]" --pre`.(close after 15 seconds)
"""
-TIPS_MAP = {
- "widgets": WIDGETS_TIPS
-}
+TIPS_MAP = {"widgets": WIDGETS_TIPS}
class TipOnStartTool:
@@ -21,7 +19,7 @@ def __init__(self, gid: str, tip_name: str):
self.gid = gid
self.slot_id = f"user-tips-{gid}"
self.tips = TIPS_MAP.get(tip_name, "")
- Thread(target=self.hide).start()
+ Thread(target=self.hide, daemon=True).start()
def show(self):
display_html(self.tips, slot_id=self.slot_id)
diff --git a/pygwalker/services/track.py b/pygwalker/services/track.py
index b8181d6d..89e204b4 100644
--- a/pygwalker/services/track.py
+++ b/pygwalker/services/track.py
@@ -1,16 +1,46 @@
from typing import Dict, Any, Optional
-import segment.analytics as analytics
-import kanaries_track
-
from pygwalker.services.global_var import GlobalVarManager
-from pygwalker.services.config import get_local_user_id
+from pygwalker.services.config import get_local_user_id, should_show_privacy_notice
+
+SEGMENT_WRITE_KEY = "z58N15R8LShkpUbBSt1ZjdDSdSEF5VpR"
+KANARIES_PUBLIC_KEY = "tk-6572d7b34a03d7fcf6cf0c86-cOzZyr6xqd"
+
+_analytics_client = None
+_kanaries_track_client = None
+
+PRIVACY_NOTICE = (
+ "PyGWalker telemetry is enabled. It only sends feature-usage events, not your analyzed data. "
+ "To opt out, run `pygwalker config --set privacy=update-only` or `pygwalker config --set privacy=offline`."
+)
+
+
+def _get_analytics_client():
+ global _analytics_client
+ if _analytics_client is None:
+ import segment.analytics as analytics
+
+ analytics.write_key = SEGMENT_WRITE_KEY
+ analytics.sync_mode = True
+ analytics.timeout = 1
+ analytics.max_retries = 0
+ _analytics_client = analytics
+ return _analytics_client
+
+
+def _get_kanaries_track_client():
+ global _kanaries_track_client
+ if _kanaries_track_client is None:
+ import kanaries_track
-analytics.write_key = 'z58N15R8LShkpUbBSt1ZjdDSdSEF5VpR'
-kanaries_public_key = "tk-6572d7b34a03d7fcf6cf0c86-cOzZyr6xqd"
-kanaries_track.config.auth_token = kanaries_public_key
-kanaries_track.config.proxies = {}
-kanaries_track.config.max_retries = 2
+ kanaries_track.config.auth_token = KANARIES_PUBLIC_KEY
+ kanaries_track.config.proxies = {}
+ kanaries_track.config.sync_send = True
+ kanaries_track.config.timeout = 1
+ kanaries_track.config.max_retries = 1
+ kanaries_track.config.thread = 0
+ _kanaries_track_client = kanaries_track
+ return _kanaries_track_client
# pylint: disable=broad-exception-caught
@@ -25,18 +55,17 @@ def track_event(event: str, properties: Optional[Dict[str, Any]] = None):
- pygwalker's mode: 'light', 'dark' or 'auto'
- pygwalker's spec type: 'json', 'file', 'url'. We won't collect the exact value of spec. No DATA YOU ANALYZE OR THEIR METADATA IS COLLECTED.
- - privacy ['offline', 'update-only', 'events'] (default: events).
+ - privacy ['offline', 'update-only', 'events'] (default: update-only).
"offline": fully offline, no data is send or api is requested
"update-only": only check whether this is a new version of pygwalker to update
"events": share which events about which feature is used in pygwalker, it only contains events data about which feature you arrive for product optimization. No DATA YOU ANALYZE IS SENT.
"""
if GlobalVarManager.privacy == "events":
try:
- analytics.track(
- user_id=get_local_user_id(),
- event=event,
- properties=properties
- )
- kanaries_track.track({**properties, "user_id": get_local_user_id()})
+ if should_show_privacy_notice():
+ print(PRIVACY_NOTICE, flush=True)
+ properties = properties or {}
+ _get_analytics_client().track(user_id=get_local_user_id(), event=event, properties=properties)
+ _get_kanaries_track_client().track({**properties, "user_id": get_local_user_id()})
except Exception:
pass
diff --git a/pygwalker/services/upload_data.py b/pygwalker/services/upload_data.py
index 691037b0..310676b4 100644
--- a/pygwalker/services/upload_data.py
+++ b/pygwalker/services/upload_data.py
@@ -12,27 +12,23 @@
def _send_js(js_code: str, slot_id: str):
display_html(
- f"""""",
- slot_id=slot_id
+ f"""""", slot_id=slot_id
)
def _send_upload_data_msg(gid: int, msg: Dict[str, Any], slot_id: str):
msg = json.dumps(msg, cls=DataFrameEncoder)
- js_code = (
- f"document.getElementById('gwalker-{gid}')?"
- ".contentWindow?"
- f".postMessage({msg}, '*');"
- )
+ js_code = f"document.getElementById('gwalker-{gid}')?.contentWindow?.postMessage({msg}, '*');"
_send_js(js_code, slot_id)
def _rand_slot_id():
- return __hash__ + '-' + rand_str(6)
+ return __hash__ + "-" + rand_str(6)
class BatchUploadDatasToolOnJupyter:
"""Upload data in batches."""
+
def run(
self,
*,
@@ -41,7 +37,7 @@ def run(
tunnel_id: str,
records: List[Dict[str, Any]],
sample_data_count: int,
- slot_count: int = 2
+ slot_count: int = 2,
) -> None:
chunk = 1 << 12
cur_slot = 0
@@ -49,12 +45,12 @@ def run(
time.sleep(1)
for i in range(sample_data_count, len(records), chunk):
- data = records[i: min(i+chunk, len(records))]
+ data = records[i : min(i + chunk, len(records))]
msg = {
- 'action': 'postData',
- 'tunnelId': tunnel_id,
- 'dataSourceId': data_source_id,
- 'data': data,
+ "action": "postData",
+ "tunnelId": tunnel_id,
+ "dataSourceId": data_source_id,
+ "data": data,
"total": len(records),
"curIndex": i,
}
@@ -63,9 +59,9 @@ def run(
cur_slot %= slot_count
finish_msg = {
- 'action': 'finishData',
- 'tunnelId': tunnel_id,
- 'dataSourceId': data_source_id,
+ "action": "finishData",
+ "tunnelId": tunnel_id,
+ "dataSourceId": data_source_id,
}
time.sleep(1)
_send_upload_data_msg(gid, finish_msg, display_slots[cur_slot])
@@ -76,30 +72,25 @@ def run(
class BatchUploadDatasToolOnWidgets:
"""Upload data in batches(use ipywidgets)"""
+
def __init__(self, comm: BaseCommunication) -> None:
self.comm = comm
- def run(
- self,
- *,
- data_source_id: str,
- records: List[Dict[str, Any]],
- sample_data_count: int
- ) -> None:
+ def run(self, *, data_source_id: str, records: List[Dict[str, Any]], sample_data_count: int) -> None:
chunk = 1 << 12
for i in range(sample_data_count, len(records), chunk):
- data = records[i: min(i+chunk, len(records))]
+ data = records[i : min(i + chunk, len(records))]
msg = {
- 'dataSourceId': data_source_id,
+ "dataSourceId": data_source_id,
"total": len(records),
"curIndex": i,
- 'data': data,
+ "data": data,
}
self.comm.send_msg_async("postData", msg)
finish_msg = {
- 'dataSourceId': data_source_id,
+ "dataSourceId": data_source_id,
}
time.sleep(0.1)
self.comm.send_msg_async("finishData", finish_msg)
diff --git a/pygwalker/spec.py b/pygwalker/spec.py
new file mode 100644
index 00000000..88fdb5c1
--- /dev/null
+++ b/pygwalker/spec.py
@@ -0,0 +1,60 @@
+import json
+import os
+from copy import deepcopy
+from typing import Any, Dict, List, Optional, Union
+
+from pygwalker.services.spec import get_spec_json
+
+
+SpecInput = Union[str, List[Any], Dict[str, Any]]
+
+
+def _json_loads(value: str) -> Any:
+ try:
+ return json.loads(value)
+ except ValueError as exc:
+ raise ValueError("spec is not a valid json") from exc
+
+
+def _is_json_string(value: str) -> bool:
+ try:
+ _json_loads(value)
+ except ValueError:
+ return False
+ return True
+
+
+def _load_spec_object(spec: SpecInput) -> Union[List[Any], Dict[str, Any]]:
+ if isinstance(spec, str):
+ if _is_json_string(spec):
+ parsed_spec = _json_loads(spec)
+ else:
+ with open(spec, "r", encoding="utf-8") as f:
+ parsed_spec = _json_loads(f.read())
+ else:
+ parsed_spec = spec
+
+ if not isinstance(parsed_spec, (dict, list)):
+ raise ValueError("spec must contain a JSON object or array")
+
+ return parsed_spec
+
+
+def migrate(spec: SpecInput, *, version: Optional[str] = None) -> Dict[str, Any]:
+ """Migrate a saved PyGWalker or Graphic Walker spec to the current schema shape."""
+ if isinstance(spec, str) and not _is_json_string(spec):
+ if spec.lstrip().startswith(("{", "[")):
+ raise ValueError("spec is not a valid json")
+ if not os.path.exists(spec):
+ raise ValueError("spec must be a dict, list, JSON string, or existing local file path")
+
+ spec_obj = _load_spec_object(spec)
+ migrated, _ = get_spec_json(deepcopy(spec_obj))
+ if version is None:
+ from pygwalker import __version__
+
+ version = __version__
+ migrated["version"] = version
+ migrated.setdefault("chart_map", {})
+ migrated.setdefault("workflow_list", [])
+ return migrated
diff --git a/pygwalker/templates/pygwalker_main_page.html b/pygwalker/templates/pygwalker_main_page.html
index edb94641..d00798e3 100644
--- a/pygwalker/templates/pygwalker_main_page.html
+++ b/pygwalker/templates/pygwalker_main_page.html
@@ -21,7 +21,14 @@
async function runGwScript() {
const script_stream = await fetch('data:application/octet-stream;base64,' + gw_script).then((res) => res.body.pipeThrough(new DecompressionStream('deflate')));
const script = await new Response(script_stream).text();
- eval(script);
+ const script_url = URL.createObjectURL(new Blob([script], {type: "text/javascript"}));
+ await new Promise((resolve, reject) => {
+ const script_element = document.createElement("script");
+ script_element.src = script_url;
+ script_element.onload = resolve;
+ script_element.onerror = reject;
+ document.body.appendChild(script_element);
+ }).finally(() => URL.revokeObjectURL(script_url));
const props_stream = await fetch('data:application/octet-stream;base64,' + props_data).then((res) => res.body.pipeThrough(new DecompressionStream('deflate')));
const props = await new Response(props_stream).json();
try{
diff --git a/pygwalker/utils/__init__.py b/pygwalker/utils/__init__.py
index aa50b388..93755f85 100644
--- a/pygwalker/utils/__init__.py
+++ b/pygwalker/utils/__init__.py
@@ -1,4 +1,3 @@
-
def fallback_value(*values):
"""Return the first non-None value in a list of values."""
for value in values:
diff --git a/pygwalker/utils/check_walker_params.py b/pygwalker/utils/check_walker_params.py
index 05057d49..606d66a9 100644
--- a/pygwalker/utils/check_walker_params.py
+++ b/pygwalker/utils/check_walker_params.py
@@ -13,6 +13,4 @@ def check_expired_params(params: Dict[str, Any]):
for old_param, new_param in expired_params_map.items():
if old_param in params:
- logger.warning(
- f"Parameter `{old_param}` is expired, please use `{new_param}` instead."
- )
+ logger.warning(f"Parameter `{old_param}` is expired, please use `{new_param}` instead.")
diff --git a/pygwalker/utils/computation.py b/pygwalker/utils/computation.py
new file mode 100644
index 00000000..345f9150
--- /dev/null
+++ b/pygwalker/utils/computation.py
@@ -0,0 +1,73 @@
+import warnings
+from typing import Optional, Tuple
+
+from pygwalker._typing import IComputation
+from pygwalker.data_parsers.database_parser import Connector
+from pygwalker.utils import fallback_value
+
+
+def _is_connector_dataset(dataset) -> bool:
+ return isinstance(dataset, (Connector, str))
+
+
+def _warn_legacy_computation_param(name: str, replacement: str) -> None:
+ warnings.warn(
+ f"`{name}` is deprecated and will be removed in a future release; use `{replacement}` instead.",
+ DeprecationWarning,
+ stacklevel=3,
+ )
+
+
+def resolve_computation_mode(
+ dataset: object,
+ *,
+ computation: Optional[IComputation] = None,
+ kernel_computation: Optional[bool] = None,
+ cloud_computation: bool = False,
+ use_kernel_calc: Optional[bool] = None,
+ default_kernel_computation: Optional[bool] = None,
+ force_kernel_for_connectors: bool = True,
+) -> Tuple[Optional[bool], bool]:
+ """Resolve public computation options to internal kernel/cloud flags."""
+ if computation is not None and computation not in ("auto", "browser", "kernel", "cloud"):
+ raise ValueError("`computation` must be one of 'auto', 'browser', 'kernel', or 'cloud'.")
+
+ explicit_auto = computation == "auto"
+ if computation is not None and not explicit_auto:
+ enabled_legacy_params = []
+ if kernel_computation is True:
+ enabled_legacy_params.append("kernel_computation")
+ if cloud_computation is True:
+ enabled_legacy_params.append("cloud_computation")
+ if use_kernel_calc is True:
+ enabled_legacy_params.append("use_kernel_calc")
+ if enabled_legacy_params:
+ legacy_names = ", ".join(enabled_legacy_params)
+ raise ValueError(
+ f"`computation` replaces legacy computation flags; remove {legacy_names} "
+ "or express the mode with `computation` only."
+ )
+
+ if computation == "browser":
+ return False, False
+ if computation == "kernel":
+ return True, False
+ if computation == "cloud":
+ return False, True
+
+ if kernel_computation is not None:
+ _warn_legacy_computation_param("kernel_computation", "computation='kernel' or computation='browser'")
+ if cloud_computation:
+ _warn_legacy_computation_param("cloud_computation", "computation='cloud'")
+ if use_kernel_calc is not None:
+ _warn_legacy_computation_param("use_kernel_calc", "computation='kernel' or computation='browser'")
+
+ if cloud_computation:
+ return False, True
+
+ use_kernel = fallback_value(kernel_computation, use_kernel_calc, default_kernel_computation)
+ if force_kernel_for_connectors and _is_connector_dataset(dataset):
+ return True, False
+ if use_kernel is None:
+ return None, False
+ return bool(use_kernel), False
diff --git a/pygwalker/utils/custom_sqlglot.py b/pygwalker/utils/custom_sqlglot.py
index e5ecf494..34b53c10 100644
--- a/pygwalker/utils/custom_sqlglot.py
+++ b/pygwalker/utils/custom_sqlglot.py
@@ -5,18 +5,12 @@
from sqlglot import exp
from sqlglot.helper import seq_get
from sqlglot.generator import Generator
-from sqlglot.dialects.dialect import (
- build_date_delta,
- build_date_delta_with_interval,
- rename_func,
- unit_to_str
-)
+from sqlglot.dialects.dialect import build_date_delta, build_date_delta_with_interval, rename_func, unit_to_str
# Duckdb Dialect
DuckdbDialect.Parser.FUNCTIONS["LOG10"] = lambda args: exp.Log(
- this=exp.Literal(this="10", is_string=False),
- expression=seq_get(args, 0)
+ this=exp.Literal(this="10", is_string=False), expression=seq_get(args, 0)
)
@@ -39,19 +33,32 @@ def _postgres_unix_to_time_sql(self: Generator, expression: exp.UnixToTime) -> s
# temporary fix for Postgres IN clause(bin filter)
def _postgres_in_sql(self: Generator, expression: exp.In) -> str:
- expression.set("expressions", [
- exp.Array(expressions=[
- exp.cast(item, to=exp.DataType.Type.DOUBLE) if isinstance(item, exp.Literal) and item.args.get("is_string") is False else item
- for item in in_item_exp.args.get("expressions", [])
- ]) if isinstance(in_item_exp, exp.Array) else in_item_exp
- for in_item_exp in expression.args.get("expressions", [])
- ])
+ expression.set(
+ "expressions",
+ [
+ exp.Array(
+ expressions=[
+ exp.cast(item, to=exp.DataType.Type.DOUBLE)
+ if isinstance(item, exp.Literal) and item.args.get("is_string") is False
+ else item
+ for item in in_item_exp.args.get("expressions", [])
+ ]
+ )
+ if isinstance(in_item_exp, exp.Array)
+ else in_item_exp
+ for in_item_exp in expression.args.get("expressions", [])
+ ],
+ )
return self.in_sql(expression)
def _postgres_timestamp_trunc(self: Generator, expression: exp.TimestampTrunc) -> str:
if expression.unit.this.lower() == "isoyear":
- return self.func("to_date", self.func("to_char", expression.this, exp.Literal.string("IYYY-0001")), exp.Literal.string("IYYY-IDDD"))
+ return self.func(
+ "to_date",
+ self.func("to_char", expression.this, exp.Literal.string("IYYY-0001")),
+ exp.Literal.string("IYYY-IDDD"),
+ )
return self.func("DATE_TRUNC", unit_to_str(expression), expression.this)
@@ -61,32 +68,47 @@ def _postgres_time_to_str_sql(self: Generator, expression: exp.TimeToStr) -> str
# postgres not support non-iso week
# current_pass_days = EXTRACT(isodow FROM DATE_TRUNC('year', date))
# week_number = floor((EXTRACT(day from date) + current_pass_days - 1) / 7)
- return self.sql(exp.Floor(
- this=exp.Div(
- this=exp.Paren(this=exp.Add(
- this=exp.Sub(
- this=exp.Cast(this=self.func("TO_CHAR", expression.this, exp.Literal.string("DDD")), to="int"),
- expression=exp.Literal.number(1)
+ return self.sql(
+ exp.Floor(
+ this=exp.Div(
+ this=exp.Paren(
+ this=exp.Add(
+ this=exp.Sub(
+ this=exp.Cast(
+ this=self.func("TO_CHAR", expression.this, exp.Literal.string("DDD")), to="int"
+ ),
+ expression=exp.Literal.number(1),
+ ),
+ expression=exp.Extract(
+ this=exp.Var(this="isodow"),
+ expression=exp.TimestampTrunc(this=expression.this, unit=exp.Literal.string("year")),
+ ),
+ )
),
- expression=exp.Extract(this=exp.Var(this="isodow"), expression=exp.TimestampTrunc(this=expression.this, unit=exp.Literal.string("year")))
- )),
- expression=exp.Literal.number(7),
+ expression=exp.Literal.number(7),
+ )
)
- ))
+ )
return self.func("TO_CHAR", expression.this, self.format_time(expression))
def _postgres_str_to_time_sql(self: Generator, expression: exp.StrToTime) -> str:
# adapter duckdb non-iso week
- if expression.args.get("format").this == "%Y%U" and isinstance(expression.this, exp.TimeToStr) and expression.this.args.get("format").this == "%Y%U":
- return self.sql(exp.Sub(
- this=exp.TimestampTrunc(this=expression.this.this, unit=exp.Literal.string("day")),
- expression=exp.Mul(
- this=exp.Extract(this=exp.Var(this="dow"), expression=expression.this.this),
- expression=exp.Interval(this=exp.Literal.number(1), unit=exp.Var(this="day"))
+ if (
+ expression.args.get("format").this == "%Y%U"
+ and isinstance(expression.this, exp.TimeToStr)
+ and expression.this.args.get("format").this == "%Y%U"
+ ):
+ return self.sql(
+ exp.Sub(
+ this=exp.TimestampTrunc(this=expression.this.this, unit=exp.Literal.string("day")),
+ expression=exp.Mul(
+ this=exp.Extract(this=exp.Var(this="dow"), expression=expression.this.this),
+ expression=exp.Interval(this=exp.Literal.number(1), unit=exp.Var(this="day")),
+ ),
)
- ))
+ )
return self.func("TO_TIMESTAMP", expression.this, self.format_time(expression))
@@ -129,12 +151,24 @@ def _mysql_extract_sql(self: Generator, expression: exp.Extract) -> str:
if unit == "week":
return self.func("WEEK", expression.expression, exp.Literal.number(3))
if unit == "isoyear":
- return self.sql(exp.Floor(this=exp.Div(this=self.func("YEARWEEK", expression.expression, exp.Literal.number(3)), expression=exp.Literal.number(100))))
+ return self.sql(
+ exp.Floor(
+ this=exp.Div(
+ this=self.func("YEARWEEK", expression.expression, exp.Literal.number(3)),
+ expression=exp.Literal.number(100),
+ )
+ )
+ )
if unit == "isodow":
- return self.sql(exp.Add(
- this=exp.Mod(this=exp.Add(this=self.func("DAYOFWEEK", expression.expression), expression=exp.Literal.number(5)), expression=exp.Literal.number(7)),
- expression=exp.Literal.number(1)
- ))
+ return self.sql(
+ exp.Add(
+ this=exp.Mod(
+ this=exp.Add(this=self.func("DAYOFWEEK", expression.expression), expression=exp.Literal.number(5)),
+ expression=exp.Literal.number(7),
+ ),
+ expression=exp.Literal.number(1),
+ )
+ )
return self.extract_sql(expression)
@@ -142,13 +176,21 @@ def _mysql_unix_to_time_sql(self: Generator, expression: exp.UnixToTime) -> str:
scale = expression.args.get("scale") or exp.UnixToTime.SECONDS
timestamp = expression.this
- return self.func("FROM_UNIXTIME", exp.Div(this=timestamp, expression=exp.func("POW", 10, scale)), self.format_time(expression))
+ return self.func(
+ "FROM_UNIXTIME", exp.Div(this=timestamp, expression=exp.func("POW", 10, scale)), self.format_time(expression)
+ )
def _mysql_str_to_time_sql(self: Generator, expression: exp.StrToTime) -> str:
# adapter duckdb non-iso week
- if expression.args.get("format").this == "%Y%U" and isinstance(expression.this, exp.TimeToStr) and expression.this.args.get("format").this == "%Y%U":
- return _mysql_timestamptrunc_sql(self, exp.TimestampTrunc(this=expression.this.this, unit=exp.Literal.string("WEEK")))
+ if (
+ expression.args.get("format").this == "%Y%U"
+ and isinstance(expression.this, exp.TimeToStr)
+ and expression.this.args.get("format").this == "%Y%U"
+ ):
+ return _mysql_timestamptrunc_sql(
+ self, exp.TimestampTrunc(this=expression.this.this, unit=exp.Literal.string("WEEK"))
+ )
return self.func("STR_TO_DATE", expression.this, self.format_time(expression))
@@ -181,11 +223,19 @@ def _snowflake_time_to_str(self: Generator, expression: exp.TimeToStr) -> str:
return self.func(
"IFF",
exp.EQ(
- this=self.func("TO_CHAR", self.func("TO_TIMESTAMP_TZ", self.func("TO_CHAR", expression.this, exp.Literal.string('YYYY')), exp.Literal.string('YYYY')), exp.Literal.string("DY")),
- expression=exp.Literal.string('Sun')
+ this=self.func(
+ "TO_CHAR",
+ self.func(
+ "TO_TIMESTAMP_TZ",
+ self.func("TO_CHAR", expression.this, exp.Literal.string("YYYY")),
+ exp.Literal.string("YYYY"),
+ ),
+ exp.Literal.string("DY"),
+ ),
+ expression=exp.Literal.string("Sun"),
),
self.func("WEEK", expression.this),
- exp.Sub(this=self.func("WEEK", expression.this), expression=exp.Literal.number(1))
+ exp.Sub(this=self.func("WEEK", expression.this), expression=exp.Literal.number(1)),
)
return self.func("TO_CHAR", exp.cast(expression.this, exp.DataType.Type.TIMESTAMP), self.format_time(expression))
@@ -193,7 +243,11 @@ def _snowflake_time_to_str(self: Generator, expression: exp.TimeToStr) -> str:
def _snowflake_str_to_time_sql(self: Generator, expression: exp.StrToTime) -> str:
# adapter duckdb non-iso week
- if expression.args.get("format").this == "%Y%U" and isinstance(expression.this, exp.TimeToStr) and expression.this.args.get("format").this == "%Y%U":
+ if (
+ expression.args.get("format").this == "%Y%U"
+ and isinstance(expression.this, exp.TimeToStr)
+ and expression.this.args.get("format").this == "%Y%U"
+ ):
return self.func("DATE_TRUNC", exp.Literal.string("WEEK"), expression.this.this)
return self.func("TO_TIMESTAMP", expression.this, self.format_time(expression))
@@ -207,9 +261,9 @@ def _snowflake_timestamp_trunc_sql(self: Generator, expression: exp.TimestampTru
exp.Var(this="day"),
exp.Sub(
this=exp.Literal.number(1),
- expression=exp.Extract(this=exp.Var(this="DAYOFWEEKISO"), expression=expression.this)
+ expression=exp.Extract(this=exp.Var(this="DAYOFWEEKISO"), expression=expression.this),
),
- self.func("date_trunc", exp.Literal.string("day"), expression.this)
+ self.func("date_trunc", exp.Literal.string("day"), expression.this),
)
# dateadd(week, 1-(WEEKISO(date)), trunc_iso_week)
@@ -217,11 +271,8 @@ def _snowflake_timestamp_trunc_sql(self: Generator, expression: exp.TimestampTru
return self.func(
"dateadd",
exp.Var(this="week"),
- exp.Sub(
- this=exp.Literal.number(1),
- expression=self.func("WEEKISO", expression.this)
- ),
- trunc_iso_week
+ exp.Sub(this=exp.Literal.number(1), expression=self.func("WEEKISO", expression.this)),
+ trunc_iso_week,
)
# duckdb week means "isoweek"
diff --git a/pygwalker/utils/dependencies.py b/pygwalker/utils/dependencies.py
new file mode 100644
index 00000000..42bfba85
--- /dev/null
+++ b/pygwalker/utils/dependencies.py
@@ -0,0 +1,8 @@
+DUCKDB_IMPORT_ERROR = (
+ "PyGWalker requires duckdb for dataframe querying. Install it with `pip install duckdb` "
+ "or reinstall PyGWalker from its project dependencies."
+)
+
+
+def raise_missing_duckdb(exc: ModuleNotFoundError) -> None:
+ raise ModuleNotFoundError(DUCKDB_IMPORT_ERROR) from exc
diff --git a/pygwalker/utils/display.py b/pygwalker/utils/display.py
index 39534b88..2c1f85c8 100644
--- a/pygwalker/utils/display.py
+++ b/pygwalker/utils/display.py
@@ -6,11 +6,7 @@
DISPLAY_HANDLER = {}
-def display_html(
- html: Union[str, HTML, ipywidgets.Widget],
- *,
- slot_id: str = None
-):
+def display_html(html: Union[str, HTML, ipywidgets.Widget], *, slot_id: str = None):
"""Judge the presentation method to be used based on the context
Args:
diff --git a/pygwalker/utils/dsl_transform.py b/pygwalker/utils/dsl_transform.py
index 18592777..49203614 100644
--- a/pygwalker/utils/dsl_transform.py
+++ b/pygwalker/utils/dsl_transform.py
@@ -1,8 +1,8 @@
-from typing import Dict, List, Any, Optional, Callable
-import os
+from typing import Dict, List, Any
+import atexit
import json
-from pygwalker._constants import ROOT_DIR
+from pygwalker.utils.frontend_assets import read_frontend_asset
from .randoms import rand_str
_dsl_to_workflow_js = None # type: Optional[Callable]
@@ -15,20 +15,38 @@
)
-def _make_js_callable(func_name, js_code):
- """Create a callable that executes a named JS function via mini-racer (V8)."""
- from py_mini_racer import MiniRacer
+class _MiniRacerCallable:
+ def __init__(self, func_name: str, js_code: str):
+ from py_mini_racer import MiniRacer
- ctx = MiniRacer()
- ctx.eval(js_code)
+ self._func_name = func_name
+ self._ctx = MiniRacer()
+ self._ctx.eval(js_code)
- def call(*args):
+ def __call__(self, *args):
if not args:
- return ctx.eval("{}()".format(func_name))
+ return self._ctx.eval("{}()".format(self._func_name))
args_json = json.dumps(args)
- return ctx.eval("{}(...{})".format(func_name, args_json))
+ return self._ctx.eval("{}(...{})".format(self._func_name, args_json))
- return call
+ def close(self):
+ if self._ctx is not None:
+ self._ctx.close()
+ self._ctx = None
+
+
+def _make_js_callable(func_name, js_code):
+ """Create a callable that executes a named JS function via mini-racer (V8)."""
+ return _MiniRacerCallable(func_name, js_code)
+
+
+def _close_js_runtime():
+ global _dsl_to_workflow_js, _vega_to_dsl_js
+ for runtime in (_dsl_to_workflow_js, _vega_to_dsl_js):
+ if runtime is not None and hasattr(runtime, "close"):
+ runtime.close()
+ _dsl_to_workflow_js = None
+ _vega_to_dsl_js = None
def _ensure_js_runtime():
@@ -38,14 +56,8 @@ def _ensure_js_runtime():
return
try:
- dsl_js_path = os.path.join(ROOT_DIR, 'templates', 'dist', 'dsl-to-workflow.umd.js')
- vega_js_path = os.path.join(ROOT_DIR, 'templates', 'dist', 'vega-to-dsl.umd.js')
-
- with open(dsl_js_path, 'r', encoding='utf8') as f:
- _dsl_to_workflow_js = _make_js_callable('main', f.read())
-
- with open(vega_js_path, 'r', encoding='utf8') as f:
- _vega_to_dsl_js = _make_js_callable('main', f.read())
+ _dsl_to_workflow_js = _make_js_callable("main", read_frontend_asset("dsl-to-workflow.umd.js"))
+ _vega_to_dsl_js = _make_js_callable("main", read_frontend_asset("vega-to-dsl.umd.js"))
except ImportError:
raise ImportError(_INSTALL_MSG)
@@ -57,9 +69,9 @@ def dsl_to_workflow(dsl: Dict[str, Any]) -> Dict[str, Any]:
def vega_to_dsl(vega_config: Dict[str, Any], fields: List[Dict[str, Any]]) -> Dict[str, Any]:
_ensure_js_runtime()
- return json.loads(_vega_to_dsl_js(json.dumps({
- "vl": vega_config,
- "allFields": fields,
- "visId": rand_str(6),
- "name": rand_str(6)
- })))
+ return json.loads(
+ _vega_to_dsl_js(json.dumps({"vl": vega_config, "allFields": fields, "visId": rand_str(6), "name": rand_str(6)}))
+ )
+
+
+atexit.register(_close_js_runtime)
diff --git a/pygwalker/utils/encode.py b/pygwalker/utils/encode.py
index 2c350e93..d6cd55e6 100644
--- a/pygwalker/utils/encode.py
+++ b/pygwalker/utils/encode.py
@@ -7,6 +7,7 @@
class DataFrameEncoder(json.JSONEncoder):
"""JSON encoder for DataFrame"""
+
def default(self, o):
if isinstance(o, datetime):
if o.tzinfo is None:
diff --git a/pygwalker/utils/estimate_tools.py b/pygwalker/utils/estimate_tools.py
index c7195324..5cc3639a 100644
--- a/pygwalker/utils/estimate_tools.py
+++ b/pygwalker/utils/estimate_tools.py
@@ -6,8 +6,11 @@
def estimate_average_data_size(datas: List[Dict[str, Any]]) -> int:
"""Estimate average data bytes size"""
- smp0 = datas[::max(len(datas)//32, 1)]
- smp1 = datas[::max(len(datas)//37, 1)]
+ if not datas:
+ return 0
+
+ smp0 = datas[:: max(len(datas) // 32, 1)]
+ smp1 = datas[:: max(len(datas) // 37, 1)]
avg_size = len(json.dumps(smp0, cls=DataFrameEncoder)) / len(smp0)
avg_size = max(avg_size, len(json.dumps(smp1, cls=DataFrameEncoder)) / len(smp1))
return avg_size
diff --git a/pygwalker/utils/free_port.py b/pygwalker/utils/free_port.py
index 5524fad5..41d3ae11 100644
--- a/pygwalker/utils/free_port.py
+++ b/pygwalker/utils/free_port.py
@@ -4,7 +4,7 @@
def find_free_port() -> int:
"""Find a free port on localhost"""
temp_socket = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
- temp_socket.bind(('localhost', 0))
+ temp_socket.bind(("localhost", 0))
_, port = temp_socket.getsockname()
temp_socket.close()
return port
diff --git a/pygwalker/utils/frontend_assets.py b/pygwalker/utils/frontend_assets.py
new file mode 100644
index 00000000..8e74f1df
--- /dev/null
+++ b/pygwalker/utils/frontend_assets.py
@@ -0,0 +1,21 @@
+import os
+import posixpath
+
+from pygwalker._constants import ROOT_DIR
+
+
+def frontend_asset_path(*path_parts: str) -> str:
+ return os.path.join(ROOT_DIR, "templates", "dist", *path_parts)
+
+
+def read_frontend_asset(*path_parts: str, encoding: str = "utf8") -> str:
+ path = frontend_asset_path(*path_parts)
+ try:
+ with open(path, "r", encoding=encoding) as f:
+ return f.read()
+ except FileNotFoundError as exc:
+ rel_path = posixpath.join("pygwalker", "templates", "dist", *path_parts)
+ raise RuntimeError(
+ f"Missing PyGWalker frontend asset: {rel_path}. "
+ "Run `./scripts/compile.sh` from the repository root before using source-checkout builds."
+ ) from exc
diff --git a/pygwalker/utils/payload_to_sql.py b/pygwalker/utils/payload_to_sql.py
index cab994ae..38f5fd95 100644
--- a/pygwalker/utils/payload_to_sql.py
+++ b/pygwalker/utils/payload_to_sql.py
@@ -1,19 +1,13 @@
from typing import Dict, List, Any
-def get_sql_from_payload(
- table_name: str,
- payload: Dict[str, Any],
- field_meta: List[Dict[str, str]] = None
-) -> str:
+def get_sql_from_payload(table_name: str, payload: Dict[str, Any], field_meta: List[Dict[str, str]] = None) -> str:
try:
from gw_dsl_parser import get_sql_from_payload as __get_sql_from_payload
except ImportError as exc:
- raise ImportError("gw_dsl_parser is not installed, please install it first. conda users please use `pip` to install it.") from exc
+ raise ImportError(
+ "gw_dsl_parser is not installed, please install it first. conda users please use `pip` to install it."
+ ) from exc
- sql = __get_sql_from_payload(
- table_name,
- payload,
- field_meta
- )
+ sql = __get_sql_from_payload(table_name, payload, field_meta)
return sql
diff --git a/pygwalker/utils/pydantic_compat.py b/pygwalker/utils/pydantic_compat.py
new file mode 100644
index 00000000..08384e83
--- /dev/null
+++ b/pygwalker/utils/pydantic_compat.py
@@ -0,0 +1,18 @@
+from typing import Any, Type, TypeVar
+
+from pydantic import BaseModel
+
+ModelT = TypeVar("ModelT", bound=BaseModel)
+PYDANTIC_V2 = hasattr(BaseModel, "model_validate")
+
+
+def model_dump(model: BaseModel, **kwargs) -> dict:
+ if PYDANTIC_V2:
+ return model.model_dump(**kwargs)
+ return model.dict(**kwargs)
+
+
+def model_validate(model_cls: Type[ModelT], data: Any) -> ModelT:
+ if PYDANTIC_V2:
+ return model_cls.model_validate(data)
+ return model_cls.parse_obj(data)
diff --git a/pygwalker/utils/randoms.py b/pygwalker/utils/randoms.py
index e5c34a0f..81fa5f20 100644
--- a/pygwalker/utils/randoms.py
+++ b/pygwalker/utils/randoms.py
@@ -4,7 +4,7 @@
def rand_str(n: int = 8, options: str = string.ascii_letters + string.digits) -> str:
- return ''.join(random.sample(options, n))
+ return "".join(random.sample(options, n))
def generate_hash_code() -> str:
diff --git a/pygwalker/utils/runtime_env.py b/pygwalker/utils/runtime_env.py
index 2b348e45..328f379d 100644
--- a/pygwalker/utils/runtime_env.py
+++ b/pygwalker/utils/runtime_env.py
@@ -4,10 +4,11 @@
def _is_jupyter() -> bool:
try:
from IPython import get_ipython
+
ip = get_ipython()
if ip is None:
return False
- return ip.has_trait('kernel')
+ return ip.has_trait("kernel")
except Exception:
return False
diff --git a/pygwalker/utils/spec.py b/pygwalker/utils/spec.py
new file mode 100644
index 00000000..514e648e
--- /dev/null
+++ b/pygwalker/utils/spec.py
@@ -0,0 +1,15 @@
+import os
+from typing import Any, Optional
+
+
+def resolve_spec_input(spec: Any, spec_path: Optional[os.PathLike[str] | str]) -> Any:
+ """Resolve explicit spec_path while keeping legacy spec behavior."""
+ if spec_path is None:
+ if isinstance(spec, os.PathLike):
+ return os.fspath(spec)
+ return spec
+
+ if spec is not None and not (isinstance(spec, str) and spec == ""):
+ raise ValueError("Pass only one of `spec` or `spec_path`; use `spec_path` for local spec files.")
+
+ return os.fspath(spec_path)
diff --git a/pygwalker_tools/metrics/api.py b/pygwalker_tools/metrics/api.py
index 69f8d4cc..613aec28 100644
--- a/pygwalker_tools/metrics/api.py
+++ b/pygwalker_tools/metrics/api.py
@@ -1,6 +1,7 @@
"""
Experimental features
"""
+
from typing import List, Dict, Any, Union, Optional
from decimal import Decimal
import json
@@ -16,12 +17,9 @@
class Chart:
"""Chart"""
+
def __init__(self, data: DataFrame, spec: Dict[str, Any]):
- self._html = to_chart_html(
- data,
- spec,
- spec_type="vega"
- )
+ self._html = to_chart_html(data, spec, spec_type="vega")
@property
def html(self) -> str:
@@ -36,6 +34,7 @@ def _repr_html_(self):
class _JSONEncoder(json.JSONEncoder):
"""JSON encoder"""
+
def default(self, o):
if isinstance(o, Decimal):
if o.is_nan():
@@ -49,7 +48,7 @@ def get_metrics_datas(
dataset: Union[DataFrame, Connector],
metrics_name: str,
field_map: Dict[str, str],
- params: Optional[Dict[str, Any]] = None
+ params: Optional[Dict[str, Any]] = None,
) -> List[Dict[str, Any]]:
"""
Example: get 1 day retention datas
@@ -111,10 +110,7 @@ def get_metrics_datas(
parser = get_parser(dataset)
sql = get_metrics_sql(
- name=metrics_name,
- field_map=field_map,
- params=params,
- origin_table_name=parser.placeholder_table_name
+ name=metrics_name, field_map=field_map, params=params, origin_table_name=parser.placeholder_table_name
)
if isinstance(dataset, Connector):
@@ -168,12 +164,13 @@ class MetricsChart:
- dimensions: ['date']
- params: ['within_active_days']
"""
+
def __init__(
self,
dataset: Union[DataFrame, Connector],
field_map: Dict[str, str],
params: Optional[Dict[str, Any]] = None,
- reverse_axis: bool = False
+ reverse_axis: bool = False,
):
self.dataset = dataset
self.field_map = field_map
@@ -182,10 +179,7 @@ def __init__(
def _get_datas(self, metrics_name: str, params: Optional[Dict[str, Any]] = None) -> pd.DataFrame:
datas = get_metrics_datas(
- dataset=self.dataset,
- metrics_name=metrics_name,
- field_map=self.field_map,
- params=params or self.params
+ dataset=self.dataset, metrics_name=metrics_name, field_map=self.field_map, params=params or self.params
)
return pd.DataFrame(json.loads(json.dumps(datas, cls=_JSONEncoder)))
@@ -201,7 +195,7 @@ def pv(self) -> Chart:
"encoding": {
"x": {"field": "date"},
"y": {"field": "pv"},
- }
+ },
}
return Chart(datas, params)
@@ -212,7 +206,7 @@ def uv(self) -> Chart:
"encoding": {
"x": {"field": "date"},
"y": {"field": "uv"},
- }
+ },
}
return Chart(datas, params)
@@ -223,7 +217,7 @@ def mau(self) -> Chart:
"encoding": {
"x": {"field": "date"},
"y": {"field": "mau"},
- }
+ },
}
return Chart(datas, params)
@@ -234,7 +228,7 @@ def retention(self) -> Chart:
"encoding": {
"x": {"field": "date"},
"y": {"field": "retention"},
- }
+ },
}
return Chart(datas, params)
@@ -245,7 +239,7 @@ def new_user_count(self) -> Chart:
"encoding": {
"x": {"field": "date"},
"y": {"field": "new_user_count"},
- }
+ },
}
return Chart(datas, params)
@@ -263,7 +257,7 @@ def cohort_matrix(self) -> Chart:
"x": {"field": "date", "type": "ordinal"},
"y": {"field": "new_user_count", "type": "ordinal"},
"color": {"field": "retention"},
- }
+ },
}
return Chart(datas, params)
@@ -274,7 +268,7 @@ def active_user_count(self) -> Chart:
"encoding": {
"x": {"field": "date"},
"y": {"field": "active_user_count"},
- }
+ },
}
return Chart(datas, params)
@@ -285,6 +279,6 @@ def user_churn_rate_base_active(self):
"encoding": {
"x": {"field": "date"},
"y": {"field": "user_churn_rate"},
- }
+ },
}
return Chart(datas, params)
diff --git a/pygwalker_tools/metrics/core.py b/pygwalker_tools/metrics/core.py
index bd92e526..e1a2108c 100644
--- a/pygwalker_tools/metrics/core.py
+++ b/pygwalker_tools/metrics/core.py
@@ -19,7 +19,7 @@
"___default_table___"
GROUP BY
strftime("date", '%Y-%m-%d')
- """
+ """,
},
"uv": {
"name": "uv",
@@ -36,7 +36,7 @@
"___default_table___"
GROUP BY
strftime("date", '%Y-%m-%d')
- """
+ """,
},
"mau": {
"name": "mau",
@@ -53,7 +53,7 @@
"___default_table___"
GROUP BY
strftime("date", '%Y-%m')
- """
+ """,
},
"retention": {
"name": "retention",
@@ -86,7 +86,7 @@
datediff('{time_unit}', t0."date", t1."date") = {time_size}
GROUP BY
strftime(t0."date", '%Y-%m-%d')
- """
+ """,
},
"new_user_count": {
"name": "new_user_count",
@@ -105,7 +105,7 @@
"date"::date = "user_signup_date"::date
GROUP BY
strftime("___default_table___"."date", '%Y-%m-%d')
- """
+ """,
},
"active_user": {
"name": "active_user",
@@ -128,7 +128,7 @@
) t1
ON
datediff('day', t1."date", t0."date") BETWEEN 0 AND {within_active_days}
- """
+ """,
},
"active_user_count": {
"name": "active_user_count",
@@ -145,7 +145,7 @@
"active_user"
GROUP BY
"active_user"."date"
- """
+ """,
},
"user_churn_rate_base_active": {
"name": "user_churn_rate_base_active",
@@ -169,15 +169,12 @@
"t0"."date"
HAVING
COUNT("t1"."user_id") > 0
- """
- }
+ """,
+ },
}
-def _replace_table_name_to_subquery(
- origin_sql: str,
- table_query_map: List[Tuple[str, str]]
-) -> str:
+def _replace_table_name_to_subquery(origin_sql: str, table_query_map: List[Tuple[str, str]]) -> str:
"""
replace table name to subquery
example:
@@ -196,22 +193,13 @@ def _replace_table_name_to_subquery(
alias_name = from_exp.this.alias
else:
alias_name = table_name
- sub_query_node = exp.Subquery(
- this=sub_query_sql_ast,
- alias=f'"{alias_name}"'
- )
+ sub_query_node = exp.Subquery(this=sub_query_sql_ast, alias=f'"{alias_name}"')
from_exp.this.replace(sub_query_node)
return origin_sql_ast.sql("duckdb")
-def get_metrics_sql(
- *,
- name: str,
- field_map: Dict[str, str],
- params: Dict[str, Any],
- origin_table_name: str
-) -> str:
+def get_metrics_sql(*, name: str, field_map: Dict[str, str], params: Dict[str, Any], origin_table_name: str) -> str:
"""get metrics sql"""
if name not in METRICS_DEFINITIONS:
raise ValueError(f"Unknown metrics name: {name}")
@@ -229,10 +217,14 @@ def get_metrics_sql(
timestamp_field = {"date"}
- field_map_sql = ",\n".join([
- f'"{field_map[field]}" "{field}"' if field not in timestamp_field else f'"{field_map[field]}"::timestamp "{field}"'
- for field in metrics_definition["fields"]
- ])
+ field_map_sql = ",\n".join(
+ [
+ f'"{field_map[field]}" "{field}"'
+ if field not in timestamp_field
+ else f'"{field_map[field]}"::timestamp "{field}"'
+ for field in metrics_definition["fields"]
+ ]
+ )
sub_query = f"""
SELECT
{field_map_sql}
@@ -241,15 +233,15 @@ def get_metrics_sql(
"""
sql = metrics_definition["sql"].format(**used_params)
- table_query_map = [
- ("___default_table___", sub_query)
- ]
+ table_query_map = [("___default_table___", sub_query)]
for depend in metrics_definition["depends"]:
- table_query_map.append((
- depend,
- get_metrics_sql(name=depend, field_map=field_map, params=params, origin_table_name=origin_table_name)
- ))
+ table_query_map.append(
+ (
+ depend,
+ get_metrics_sql(name=depend, field_map=field_map, params=params, origin_table_name=origin_table_name),
+ )
+ )
sql = _replace_table_name_to_subquery(sql, table_query_map)
diff --git a/pyproject.toml b/pyproject.toml
index 8267b782..336af495 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -1,7 +1,7 @@
[project]
name = "pygwalker"
dynamic = ["version"]
-requires-python = ">=3.7"
+requires-python = ">=3.10"
description = "pygwalker: turn your data into an interactive UI for data exploration and visualization"
authors = [ { name = "kanaries", email = "support@kanaries.net" } ]
license-files = { paths = ["LICENSE"] }
@@ -16,15 +16,15 @@ dependencies = [
"ipython",
"astor",
"typing_extensions",
- "ipywidgets",
- "pydantic",
+ "ipywidgets>=7.0.0",
+ "pydantic>=1.10,<3",
"psutil",
"duckdb>=0.10.4,<2.0.0",
- "pyarrow",
+ "pyarrow>=10,<25",
"sqlglot>=23.15.8",
- "requests",
+ "requests>=2.31,<3",
"arrow",
- "sqlalchemy",
+ "sqlalchemy>=1.4,<3",
"gw_dsl_parser==0.1.49.1",
"appdirs",
"segment-analytics-python==2.2.3",
@@ -49,6 +49,16 @@ polars = ["polars"]
streamlit = ["streamlit"]
reflex = ["reflex"]
notebook = [
+ "jupyter-client>7.4.9",
+ "jupyter-server>2.5.0",
+ "ipywidgets>=8.0.0"
+]
+labv4 = [
+ "jupyter-client>7.4.9",
+ "jupyter-server>2.5.0",
+ "ipywidgets>=8.0.0"
+]
+notebook-legacy = [
"jupyter-client<=7.4.9,>6.0.0",
"jupyter-server<=2.5.0",
"ipywidgets<8.0.0,>7.0.0"
@@ -59,17 +69,15 @@ snowflake = [
"snowflake-sqlalchemy==1.5.0",
"pyarrow==10.0.1"
]
-labv4 = [
- "jupyter-client>7.4.9",
- "jupyter-server>2.5.0",
- "ipywidgets>=8.0.0"
-]
export = ["mini-racer>=0.12"]
all = [
"pygwalker[pandas,polars,streamlit,reflex,export]",
]
dev = [
"build",
+ "pytest",
+ "nbmake",
+ "ruff",
"twine",
"jupyterlab",
"jupyter_server_proxy",
@@ -97,7 +105,6 @@ artifacts = [ "pygwalker/templates/*" ]
dependencies = ["hatch-jupyter-builder"]
build-function = "hatch_jupyter_builder.npm_builder"
ensured-targets = ["pygwalker/templates/dist/pygwalker-app.iife.js"]
-skip-if-exists = ["pygwalker/templates/dist/pygwalker-app.iife.js"]
# install-pre-commit-hook = true
optional-editable-build = true
@@ -128,3 +135,18 @@ include = [
[build-system]
requires = ["hatchling"]
build-backend = "hatchling.build"
+
+[tool.ruff]
+target-version = "py310"
+line-length = 120
+extend-exclude = [
+ "*.ipynb",
+ "app",
+ "pygwalker/templates/dist",
+]
+
+[tool.ruff.lint]
+select = ["E4", "E7", "E9", "F"]
+
+[tool.ruff.lint.per-file-ignores]
+"pygwalker/__init__.py" = ["E402"]
diff --git a/scripts/ci_run_pytest.py b/scripts/ci_run_pytest.py
deleted file mode 100644
index dc70fc1e..00000000
--- a/scripts/ci_run_pytest.py
+++ /dev/null
@@ -1,100 +0,0 @@
-#!/usr/bin/env python3
-import subprocess
-import sys
-import time
-import xml.etree.ElementTree as ET
-from pathlib import Path
-
-JUNIT_PATH = Path("pytest-results.xml")
-HARD_TIMEOUT_SECONDS = 20 * 60
-PASS_GRACE_SECONDS = 30
-POLL_SECONDS = 1
-
-
-def _terminate_process(proc: subprocess.Popen) -> None:
- proc.terminate()
- try:
- proc.wait(timeout=10)
- except subprocess.TimeoutExpired:
- proc.kill()
- proc.wait(timeout=5)
-
-
-def _junit_all_passed(junit_path: Path) -> bool:
- if not junit_path.exists():
- return False
-
- try:
- root = ET.parse(junit_path).getroot()
- except ET.ParseError:
- return False
-
- if root.tag == "testsuite":
- suites = [root]
- else:
- suites = root.findall("testsuite")
-
- if not suites:
- return False
-
- total_tests = 0
- total_failures = 0
- total_errors = 0
- for suite in suites:
- total_tests += int(suite.attrib.get("tests", 0))
- total_failures += int(suite.attrib.get("failures", 0))
- total_errors += int(suite.attrib.get("errors", 0))
-
- return total_tests > 0 and total_failures == 0 and total_errors == 0
-
-
-def main() -> int:
- cmd = [
- sys.executable,
- "-X",
- "faulthandler",
- "-m",
- "pytest",
- "-o",
- "faulthandler_timeout=300",
- "--junitxml=pytest-results.xml",
- "tests",
- ]
-
- print("[CI] pytest start", flush=True)
- proc = subprocess.Popen(cmd)
- start_time = time.time()
- pass_detected_at = None
-
- while True:
- rc = proc.poll()
- if rc is not None:
- print("[CI] pytest end", flush=True)
- return rc
-
- if _junit_all_passed(JUNIT_PATH):
- if pass_detected_at is None:
- pass_detected_at = time.time()
- elif time.time() - pass_detected_at >= PASS_GRACE_SECONDS:
- print(
- "[CI] pytest summary indicates success but process is still alive; "
- "terminating stuck process and continuing.",
- flush=True
- )
- _terminate_process(proc)
- print("[CI] pytest end", flush=True)
- return 0
- else:
- pass_detected_at = None
-
- if time.time() - start_time >= HARD_TIMEOUT_SECONDS:
- print("[CI] pytest watchdog timeout reached; terminating process.", flush=True)
- _terminate_process(proc)
- print("[CI] pytest end", flush=True)
- return 1
-
- time.sleep(POLL_SECONDS)
-
-
-if __name__ == "__main__":
- raise SystemExit(main())
diff --git a/scripts/compile.sh b/scripts/compile.sh
index 72ec24ed..3558673b 100755
--- a/scripts/compile.sh
+++ b/scripts/compile.sh
@@ -10,4 +10,4 @@ APP=$file_dir/app
ret=$?
cd $cur_dir
-exit $?
\ No newline at end of file
+exit $ret
diff --git a/scripts/local_ci.py b/scripts/local_ci.py
new file mode 100644
index 00000000..4253a6e8
--- /dev/null
+++ b/scripts/local_ci.py
@@ -0,0 +1,84 @@
+"""Run the local equivalent of the GitHub Actions Auto CI workflow."""
+
+from __future__ import annotations
+
+import argparse
+import subprocess
+import sys
+from pathlib import Path
+
+
+REPO_ROOT = Path(__file__).resolve().parents[1]
+APP_DIR = REPO_ROOT / "app"
+TESTS_DIR = REPO_ROOT / "tests"
+PYTHON_TARGETS = ["pygwalker", "tests", "scripts", "bin", "pygwalker_tools"]
+
+
+def run(command: list[str], cwd: Path = REPO_ROOT) -> None:
+ print(f"\n$ {' '.join(command)}", flush=True)
+ subprocess.run(command, cwd=cwd, check=True)
+
+
+def run_frontend_ci() -> None:
+ run(["sh", "scripts/compile.sh"])
+ run(["yarn", "playwright", "install", "--with-deps", "chromium"], cwd=APP_DIR)
+ run(["yarn", "test:front_end"], cwd=APP_DIR)
+
+
+def run_notebook_ci() -> None:
+ run([sys.executable, "scripts/test-init.py"])
+ run([sys.executable, "-m", "pip", "install", "ipykernel", "nbmake", "pandas", "polars", "pytest"])
+ run([sys.executable, "-m", "ipykernel", "install", "--name", "python", "--user"])
+ run(["jupyter", "kernelspec", "list"])
+
+ notebooks = sorted(path.name for path in TESTS_DIR.glob("*.ipynb"))
+ run([sys.executable, "-m", "pytest", "--nbmake", "--nbmake-kernel=python", *notebooks], cwd=TESTS_DIR)
+
+
+def install_legacy_modin_ci_deps() -> None:
+ run([sys.executable, "-m", "pip", "install", "numpy<2", "pandas<2.1"])
+ run([sys.executable, "-m", "pip", "install", "aiohttp==3.8.6"])
+ run([sys.executable, "-m", "pip", "install", "modin==0.23.1", "modin[ray]==0.23.1"])
+ run([sys.executable, "-m", "pip", "install", "pydantic==1.10.9"])
+
+
+def run_python_ci() -> None:
+ run([sys.executable, "-m", "pip", "install", "duckdb_engine"])
+ run([sys.executable, "-m", "pip", "install", "pytest", "ruff", "starlette", "polars", "tornado"])
+ run([sys.executable, "-m", "ruff", "check", *PYTHON_TARGETS])
+ run([sys.executable, "-m", "ruff", "format", "--check", *PYTHON_TARGETS])
+ run(
+ [
+ sys.executable,
+ "-X",
+ "faulthandler",
+ "-W",
+ "error::DeprecationWarning:pygwalker",
+ "-m",
+ "pytest",
+ "-o",
+ "faulthandler_timeout=300",
+ "tests",
+ ]
+ )
+
+
+def main() -> int:
+ parser = argparse.ArgumentParser(description=__doc__)
+ parser.add_argument("--skip-frontend", action="store_true", help="Skip build-js equivalent checks.")
+ parser.add_argument("--skip-notebooks", action="store_true", help="Skip nbmake notebook execution.")
+ parser.add_argument("--legacy-modin-deps", action="store_true", help="Mirror the Ubuntu Python 3.11 modin CI leg.")
+ args = parser.parse_args()
+
+ if not args.skip_frontend:
+ run_frontend_ci()
+ if args.legacy_modin_deps:
+ install_legacy_modin_ci_deps()
+ if not args.skip_notebooks:
+ run_notebook_ci()
+ run_python_ci()
+ return 0
+
+
+if __name__ == "__main__":
+ raise SystemExit(main())
diff --git a/scripts/test-init.py b/scripts/test-init.py
index 093d91d5..8cb2f4fd 100644
--- a/scripts/test-init.py
+++ b/scripts/test-init.py
@@ -16,4 +16,4 @@
urlretrieve(url, f"{dir}/bike_sharing_dc.csv")
except Exception as e:
print(f"Error: could not download the file\n{e}", file=sys.stderr)
- sys.exit(1)
\ No newline at end of file
+ sys.exit(1)
diff --git a/tests/test_check_update.py b/tests/test_check_update.py
new file mode 100644
index 00000000..9564839d
--- /dev/null
+++ b/tests/test_check_update.py
@@ -0,0 +1,83 @@
+from pygwalker.services import check_update
+from pygwalker.services.global_var import GlobalVarManager
+
+
+def test_check_update_starts_daemon_thread(monkeypatch):
+ started_threads = []
+ previous_privacy = GlobalVarManager.privacy
+ GlobalVarManager.privacy = "events"
+
+ class FakeThread:
+ def __init__(self, *, target, daemon):
+ self.target = target
+ self.daemon = daemon
+
+ def start(self):
+ started_threads.append(self)
+
+ monkeypatch.setattr(check_update, "Thread", FakeThread)
+
+ try:
+ check_update.check_update()
+ finally:
+ GlobalVarManager.privacy = previous_privacy
+
+ assert len(started_threads) == 1
+ assert started_threads[0].target is check_update._check_update
+ assert started_threads[0].daemon is True
+
+
+def test_check_update_starts_thread_in_update_only_mode(monkeypatch):
+ started_threads = []
+ previous_privacy = GlobalVarManager.privacy
+ GlobalVarManager.privacy = "update-only"
+
+ class FakeThread:
+ def __init__(self, *, target, daemon):
+ self.target = target
+ self.daemon = daemon
+
+ def start(self):
+ started_threads.append(self)
+
+ monkeypatch.setattr(check_update, "Thread", FakeThread)
+
+ try:
+ check_update.check_update()
+ finally:
+ GlobalVarManager.privacy = previous_privacy
+
+ assert len(started_threads) == 1
+ assert started_threads[0].target is check_update._check_update
+ assert started_threads[0].daemon is True
+
+
+def test_check_update_does_not_start_thread_in_offline_mode(monkeypatch):
+ started_threads = []
+ previous_privacy = GlobalVarManager.privacy
+ GlobalVarManager.privacy = "offline"
+
+ class FakeThread:
+ def __init__(self, *args, **kwargs):
+ pass
+
+ def start(self):
+ started_threads.append(self)
+
+ monkeypatch.setattr(check_update, "Thread", FakeThread)
+
+ try:
+ check_update.check_update()
+ finally:
+ GlobalVarManager.privacy = previous_privacy
+
+ assert started_threads == []
+
+
+def test_check_update_suppresses_background_request_errors(monkeypatch):
+ def fail_request(_url):
+ raise OSError("network unavailable")
+
+ monkeypatch.setattr(check_update, "_request_on_python", fail_request)
+
+ assert check_update._check_update() == {}
diff --git a/tests/test_ci_workflow.py b/tests/test_ci_workflow.py
new file mode 100644
index 00000000..1e03db29
--- /dev/null
+++ b/tests/test_ci_workflow.py
@@ -0,0 +1,71 @@
+from pathlib import Path
+
+
+REPO_ROOT = Path(__file__).resolve().parents[1]
+
+
+def test_auto_ci_fails_on_pygwalker_deprecation_warnings():
+ workflow = (REPO_ROOT / ".github/workflows/auto-ci.yml").read_text(encoding="utf-8")
+
+ assert "-W error::DeprecationWarning:pygwalker" in workflow
+
+
+def test_auto_ci_enforces_python_and_frontend_quality_gates():
+ workflow = (REPO_ROOT / ".github/workflows/auto-ci.yml").read_text(encoding="utf-8")
+
+ assert "python-version: ['3.10', '3.11', '3.12', '3.13']" in workflow
+ assert 'pip install "numpy<2" "pandas<2.1"' in workflow
+ assert "ruff check pygwalker tests scripts bin pygwalker_tools" in workflow
+ assert "ruff format --check pygwalker tests scripts bin pygwalker_tools" in workflow
+ assert "pip install pytest ruff starlette polars tornado" in workflow
+ assert "yarn playwright install --with-deps chromium" in workflow
+ assert "yarn test:front_end" in workflow
+
+
+def test_local_ci_script_mirrors_auto_ci_quality_gates():
+ script = (REPO_ROOT / "scripts/local_ci.py").read_text(encoding="utf-8")
+
+ assert '"scripts/compile.sh"' in script
+ assert '"yarn", "playwright", "install", "--with-deps", "chromium"' in script
+ assert '"yarn", "test:front_end"' in script
+ assert '"numpy<2", "pandas<2.1"' in script
+ assert '"pytest", "ruff", "starlette", "polars", "tornado"' in script
+ assert '"--nbmake", "--nbmake-kernel=python"' in script
+ assert '"ruff", "check", *PYTHON_TARGETS' in script
+ assert '"ruff", "format", "--check", *PYTHON_TARGETS' in script
+ assert '"error::DeprecationWarning:pygwalker"' in script
+ assert '"faulthandler_timeout=300"' in script
+
+
+def test_auto_ci_runs_pytest_directly_without_watchdog():
+ workflow = (REPO_ROOT / ".github/workflows/auto-ci.yml").read_text(encoding="utf-8")
+
+ assert "scripts/ci_run_pytest.py" not in workflow
+ assert not (REPO_ROOT / "scripts/ci_run_pytest.py").exists()
+
+
+def test_auto_ci_runs_notebooks_through_pytest_nbmake():
+ workflow = (REPO_ROOT / ".github/workflows/auto-ci.yml").read_text(encoding="utf-8")
+
+ assert "Path('.').glob('*.ipynb')" in workflow
+ assert "python -m pytest --nbmake --nbmake-kernel=python *.ipynb" not in workflow
+ assert "jupyter nbconvert --execute" not in workflow
+
+
+def test_publish_workflow_packages_fresh_frontend_bundle():
+ workflow = (REPO_ROOT / ".github/workflows/publish.yml").read_text(encoding="utf-8")
+
+ assert "node-version: [22.x]" in workflow
+ assert "./scripts/compile.sh" in workflow
+ assert "name: pygwalker-app" in workflow
+ assert "path: ./pygwalker/templates/dist/*" in workflow
+ assert "build-py:\n needs: [build-js]" in workflow
+ assert "actions/download-artifact@v4" in workflow
+ assert "path: ./pygwalker/templates/dist" in workflow
+ assert "python -m build ." in workflow
+ assert "pypa/gh-action-pypi-publish" in workflow
+
+ build_py_start = workflow.index(" build-py:")
+ download_dist_start = workflow.index("uses: actions/download-artifact@v4", build_py_start)
+ build_package_start = workflow.index("python -m build .", build_py_start)
+ assert download_dist_start < build_package_start
diff --git a/tests/test_cloud_communication.py b/tests/test_cloud_communication.py
new file mode 100644
index 00000000..c06fab6c
--- /dev/null
+++ b/tests/test_cloud_communication.py
@@ -0,0 +1,125 @@
+import json
+from types import SimpleNamespace
+
+from pygwalker import __version__
+from pygwalker.communications.protocol import (
+ AskSpecRequest,
+ ChatChartRequest,
+ UploadCloudChartRequest,
+ UploadCloudDashboardRequest,
+ UploadSpecToCloudRequest,
+)
+from pygwalker.services import cloud_communication as cloud_communication_module
+from pygwalker.services.cloud_communication import CloudCommunicationService
+from pygwalker.services.global_var import GlobalVarManager
+
+
+def test_cloud_communication_upload_spec_updates_token_and_writes_workspace(monkeypatch):
+ previous_token = GlobalVarManager.kanaries_api_key
+ config_updates = []
+ writes = []
+ walker = SimpleNamespace(
+ spec_manager=SimpleNamespace(build_spec_obj=lambda version: {"version": version, "config": []}),
+ cloud_service=SimpleNamespace(
+ get_kanaries_user_info=lambda: {"workspaceName": "workspace"},
+ write_config_to_cloud=lambda path, data: writes.append((path, json.loads(data))),
+ ),
+ )
+ monkeypatch.setattr(cloud_communication_module, "set_config", lambda config: config_updates.append(config))
+
+ try:
+ response = CloudCommunicationService(walker).upload_spec_to_cloud(
+ UploadSpecToCloudRequest(fileName="chart.json", newToken="new-token")
+ )
+ finally:
+ GlobalVarManager.kanaries_api_key = previous_token
+
+ assert response == {"specFilePath": "workspace/chart.json"}
+ assert config_updates == [{"kanaries_token": "new-token"}]
+ assert writes == [("workspace/chart.json", {"version": __version__, "config": []})]
+
+
+def test_cloud_communication_uses_custom_text_callbacks():
+ ask_calls = []
+ chat_calls = []
+ walker = SimpleNamespace(
+ other_props={
+ "custom_ask_callback": lambda metas, query: ask_calls.append((metas, query)) or {"chart": "bar"},
+ "custom_chat_callback": lambda metas, chats: chat_calls.append((metas, chats)) or {"chart": "line"},
+ },
+ cloud_service=SimpleNamespace(),
+ )
+ service = CloudCommunicationService(walker)
+
+ ask_response = service.get_spec_by_text(AskSpecRequest(metas=[{"fid": "city"}], query="show city"))
+ chat_response = service.get_chart_by_chats(
+ ChatChartRequest(metas=[{"fid": "city"}], chats=[{"role": "user", "content": "show city"}])
+ )
+
+ assert ask_response == {"data": {"chart": "bar"}}
+ assert chat_response == {"data": {"chart": "line"}}
+ assert ask_calls == [([{"fid": "city"}], "show city")]
+ assert chat_calls == [([{"fid": "city"}], [{"role": "user", "content": "show city"}])]
+
+
+def test_cloud_communication_uploads_chart_and_dashboard():
+ upload_chart_calls = []
+ upload_dashboard_calls = []
+ walker = SimpleNamespace(
+ data_parser=object(),
+ appearance="light",
+ cloud_service=SimpleNamespace(
+ upload_cloud_chart=lambda **kwargs: (
+ upload_chart_calls.append(kwargs) or {"chart_id": "chart-id", "dataset_id": "dataset-id"}
+ ),
+ upload_cloud_dashboard=lambda **kwargs: (
+ upload_dashboard_calls.append(kwargs) or {"dashboard_id": "dashboard-id", "dataset_id": "dataset-id"}
+ ),
+ ),
+ )
+ service = CloudCommunicationService(walker)
+
+ chart_response = service.upload_to_cloud_charts(
+ UploadCloudChartRequest(
+ chartName="Chart",
+ datasetName="Dataset",
+ isPublic=True,
+ visSpec=[{"name": "Chart"}],
+ workflow=[{"type": "view"}],
+ )
+ )
+ dashboard_response = service.upload_to_cloud_dashboard(
+ UploadCloudDashboardRequest(
+ chartName="Dashboard",
+ datasetName="Dataset",
+ isPublic=False,
+ isCreateDashboard=True,
+ visSpec=[{"name": "Chart"}],
+ workflowList=[[{"type": "view"}]],
+ )
+ )
+
+ assert chart_response == {"chartId": "chart-id", "datasetId": "dataset-id"}
+ assert dashboard_response == {"dashboardId": "dashboard-id", "datasetId": "dataset-id"}
+ assert upload_chart_calls == [
+ {
+ "data_parser": walker.data_parser,
+ "chart_name": "Chart",
+ "dataset_name": "Dataset",
+ "workflow": [{"type": "view"}],
+ "spec_list": [{"name": "Chart"}],
+ "is_public": True,
+ }
+ ]
+ assert upload_dashboard_calls == [
+ {
+ "data_parser": walker.data_parser,
+ "dashboard_name": "Dashboard",
+ "dataset_name": "Dataset",
+ "workflow_list": [[{"type": "view"}]],
+ "spec_list": [{"name": "Chart"}],
+ "is_public": False,
+ "create_dashboard_flag": True,
+ "appearance": "light",
+ }
+ ]
diff --git a/tests/test_communication_transports.py b/tests/test_communication_transports.py
new file mode 100644
index 00000000..8e41f0b2
--- /dev/null
+++ b/tests/test_communication_transports.py
@@ -0,0 +1,100 @@
+import asyncio
+import json
+
+from pygwalker.communications import gradio_comm
+from pygwalker.communications.anywidget_comm import AnywidgetCommunication
+from pygwalker.communications.base import BaseCommunication
+from pygwalker.communications.hacker_comm import HackerCommunication
+from pygwalker.errors import ErrorCode
+
+
+class _FakeWidget:
+ def __init__(self):
+ self.sent = []
+
+ def send(self, message):
+ self.sent.append(message)
+
+
+def _decode_widget_response(widget):
+ assert widget.sent[0]["type"] == "pyg_response"
+ return json.loads(widget.sent[0]["data"])
+
+
+def test_anywidget_transport_returns_protocol_error_for_missing_action():
+ comm = AnywidgetCommunication("widget-gid")
+ widget = _FakeWidget()
+ comm.widget = widget
+
+ comm._on_mesage(None, {"type": "pyg_request", "msg": {"rid": "request-1", "data": {}}}, [])
+
+ message = _decode_widget_response(widget)
+ assert message["action"] == "finish_request"
+ assert message["rid"] == "request-1"
+ assert message["data"]["code"] == ErrorCode.INVALID_REQUEST
+ assert "action" in message["data"]["message"]
+
+
+def test_hacker_transport_returns_protocol_error_for_missing_action():
+ comm = HackerCommunication.__new__(HackerCommunication)
+ BaseCommunication.__init__(comm, "hacker-gid")
+ comm._HackerCommunication__increase = 0
+ sent = []
+ comm.send_msg_async = lambda action, data, rid=None: sent.append({"action": action, "data": data, "rid": rid})
+
+ comm._on_mesage({"new": json.dumps({"rid": "request-1", "data": {}})})
+
+ assert sent == [
+ {
+ "action": "finish_request",
+ "data": {
+ "code": ErrorCode.INVALID_REQUEST,
+ "data": None,
+ "message": sent[0]["data"]["message"],
+ },
+ "rid": "request-1",
+ }
+ ]
+ assert "action" in sent[0]["data"]["message"]
+
+
+class _FakeRequest:
+ def __init__(self, gid, payload):
+ self.path_params = {"gid": gid}
+ self._payload = payload
+
+ async def json(self):
+ return self._payload
+
+
+def test_gradio_router_returns_protocol_error_for_missing_action():
+ gradio_comm.GradioCommunication("gradio-gid")
+ try:
+ response = asyncio.run(gradio_comm._pygwalker_router(_FakeRequest("gradio-gid", {"data": {}})))
+ finally:
+ gradio_comm.gradio_comm_map.pop("gradio-gid", None)
+
+ payload = json.loads(response.body)
+ assert payload["code"] == ErrorCode.INVALID_REQUEST
+ assert payload["data"] is None
+ assert "action" in payload["message"]
+
+
+def test_gradio_router_preserves_successful_envelope_routing():
+ comm = gradio_comm.GradioCommunication("gradio-gid")
+ comm.register("ping", lambda _: {})
+ try:
+ response = asyncio.run(gradio_comm._pygwalker_router(_FakeRequest("gradio-gid", {"action": "ping"})))
+ finally:
+ gradio_comm.gradio_comm_map.pop("gradio-gid", None)
+
+ assert json.loads(response.body) == {"code": 0, "data": {}, "message": "success"}
+
+
+def test_comm_envelope_accepts_integer_gid_before_dispatch():
+ comm = BaseCommunication("123")
+ comm.register("ping", lambda _: {})
+
+ response = comm._receive_msg_envelope({"gid": 123, "rid": "request-1", "action": "ping", "data": {}})
+
+ assert response == {"code": 0, "data": {}, "message": "success"}
diff --git a/tests/test_computation.py b/tests/test_computation.py
new file mode 100644
index 00000000..41827df1
--- /dev/null
+++ b/tests/test_computation.py
@@ -0,0 +1,39 @@
+import pytest
+
+from pygwalker.utils.computation import resolve_computation_mode
+
+
+@pytest.mark.parametrize(
+ ("computation", "expected"),
+ [
+ ("auto", (None, False)),
+ ("browser", (False, False)),
+ ("kernel", (True, False)),
+ ("cloud", (False, True)),
+ ],
+)
+def test_resolve_computation_mode_from_new_api(computation, expected):
+ assert resolve_computation_mode(object(), computation=computation) == expected
+
+
+def test_resolve_computation_mode_preserves_unspecified_auto_detection():
+ assert resolve_computation_mode(object()) == (None, False)
+
+
+def test_resolve_computation_mode_rejects_invalid_mode():
+ with pytest.raises(ValueError, match="`computation` must be one of"):
+ resolve_computation_mode(object(), computation="remote")
+
+
+def test_resolve_computation_mode_rejects_enabled_legacy_flags_with_new_api():
+ with pytest.raises(ValueError, match="replaces legacy computation flags"):
+ resolve_computation_mode(object(), computation="browser", kernel_computation=True)
+
+
+def test_resolve_computation_mode_preserves_default_kernel_semantics():
+ assert resolve_computation_mode(object(), default_kernel_computation=True) == (True, False)
+
+
+def test_resolve_computation_mode_warns_for_legacy_kernel_param():
+ with pytest.warns(DeprecationWarning, match="kernel_computation"):
+ assert resolve_computation_mode(object(), kernel_computation=True) == (True, False)
diff --git a/tests/test_data_bridge.py b/tests/test_data_bridge.py
new file mode 100644
index 00000000..b80e6848
--- /dev/null
+++ b/tests/test_data_bridge.py
@@ -0,0 +1,166 @@
+from types import SimpleNamespace
+
+import pandas as pd
+import pyarrow as pa
+
+from pygwalker.services import data_bridge as data_bridge_module
+from pygwalker.services.data_bridge import DataBridge
+from pygwalker._constants import JUPYTER_BYTE_LIMIT
+
+
+class FakeCloudService:
+ def __init__(self):
+ self.created = []
+
+ def create_cloud_dataset(self, data_parser, name, is_public, temporary):
+ self.created.append(
+ {
+ "data_parser": data_parser,
+ "name": name,
+ "is_public": is_public,
+ "temporary": temporary,
+ }
+ )
+ return "cloud-dataset-id"
+
+
+def test_data_bridge_initializes_browser_dataframe_path():
+ bridge = DataBridge(
+ dataset=pd.DataFrame([{"city": "London", "value": 1}]),
+ field_specs=[],
+ cloud_computation=False,
+ kernel_computation=False,
+ kanaries_api_key="",
+ cloud_service=FakeCloudService(),
+ )
+
+ assert bridge.kernel_computation is False
+ assert bridge.origin_data_source == [{"city": "London", "value": 1}]
+ assert bridge.dataset_type == "pandas_dataframe"
+ assert bridge.parse_dsl_type == "client"
+ assert [field["fid"] for field in bridge.field_specs] == ["city", "value"]
+
+
+def test_data_bridge_initializes_pyarrow_table_path():
+ bridge = DataBridge(
+ dataset=pa.table({"city": ["London"], "value": [1]}),
+ field_specs=[],
+ cloud_computation=False,
+ kernel_computation=False,
+ kanaries_api_key="",
+ cloud_service=FakeCloudService(),
+ )
+
+ assert bridge.kernel_computation is False
+ assert bridge.origin_data_source == [{"city": "London", "value": 1}]
+ assert bridge.dataset_type == "pyarrow_table"
+ assert bridge.parse_dsl_type == "client"
+ assert [field["fid"] for field in bridge.field_specs] == ["city", "value"]
+
+
+def test_data_bridge_initializes_empty_dataframe_path():
+ bridge = DataBridge(
+ dataset=pd.DataFrame({"city": pd.Series(dtype="object"), "value": pd.Series(dtype="int64")}),
+ field_specs=[],
+ cloud_computation=False,
+ kernel_computation=None,
+ kanaries_api_key="",
+ cloud_service=FakeCloudService(),
+ )
+
+ assert bridge.kernel_computation is False
+ assert bridge.origin_data_source == []
+ assert bridge.dataset_type == "pandas_dataframe"
+ assert bridge.parse_dsl_type == "client"
+ assert bridge.field_specs == [
+ {"fid": "city", "name": "city", "semanticType": "nominal", "analyticType": "dimension"},
+ {"fid": "value", "name": "value", "semanticType": "quantitative", "analyticType": "dimension"},
+ ]
+
+
+def test_data_bridge_initializes_kernel_sample_path():
+ bridge = DataBridge(
+ dataset=pd.DataFrame([{"city": "London", "value": 1}]),
+ field_specs=[],
+ cloud_computation=False,
+ kernel_computation=True,
+ kanaries_api_key="",
+ cloud_service=FakeCloudService(),
+ )
+
+ assert bridge.kernel_computation is True
+ assert bridge.origin_data_source == [{"city": "London", "value": 1}]
+
+
+def test_data_bridge_auto_enables_kernel_for_large_data(monkeypatch):
+ calls = []
+
+ class FakeParser:
+ data_size = JUPYTER_BYTE_LIMIT + 1
+ raw_fields = [{"fid": "city"}]
+ dataset_type = "pandas_dataframe"
+
+ def to_records(self, limit=None):
+ calls.append(limit)
+ return [{"city": "London"}]
+
+ monkeypatch.setattr(data_bridge_module, "get_parser", lambda *_args, **_kwargs: FakeParser())
+
+ bridge = DataBridge(
+ dataset=object(),
+ field_specs=[],
+ cloud_computation=False,
+ kernel_computation=None,
+ kanaries_api_key="",
+ cloud_service=FakeCloudService(),
+ )
+
+ assert bridge.kernel_computation is True
+ assert calls == [500]
+ assert bridge.origin_data_source == [{"city": "London"}]
+
+
+def test_data_bridge_routes_cloud_computation_through_cloud_dataset(monkeypatch):
+ calls = []
+
+ class FakeParser:
+ data_size = 1
+ raw_fields = [{"fid": "city"}]
+ dataset_type = "pandas_dataframe"
+
+ def to_records(self, limit=None):
+ return [{"city": "London"}]
+
+ class FakeCloudParser(FakeParser):
+ dataset_type = "cloud_dataset"
+
+ def fake_get_parser(dataset, field_specs, other_params):
+ calls.append((dataset, field_specs, other_params))
+ if dataset == "cloud-dataset-id":
+ return FakeCloudParser()
+ return FakeParser()
+
+ monkeypatch.setattr(data_bridge_module, "get_parser", fake_get_parser)
+ cloud_service = FakeCloudService()
+
+ bridge = DataBridge(
+ dataset=object(),
+ field_specs=[],
+ cloud_computation=True,
+ kernel_computation=None,
+ kanaries_api_key="secret",
+ cloud_service=cloud_service,
+ )
+
+ assert calls[0][0] != "cloud-dataset-id"
+ assert calls[0][2] == {"kanaries_api_key": "secret"}
+ assert calls[1][0] == "cloud-dataset-id"
+ assert cloud_service.created[0]["temporary"] is True
+ assert bridge.dataset_type == "cloud_dataset"
+ assert bridge.parse_dsl_type == "server"
+
+
+def test_data_bridge_parse_dsl_type_tracks_dataset_location():
+ assert DataBridge.get_parse_dsl_type(SimpleNamespace(dataset_type="connector_database")) == "server"
+ assert DataBridge.get_parse_dsl_type(SimpleNamespace(dataset_type="cloud_dataset")) == "server"
+ assert DataBridge.get_parse_dsl_type(SimpleNamespace(dataset_type="pandas_dataframe")) == "client"
diff --git a/tests/test_data_communication.py b/tests/test_data_communication.py
new file mode 100644
index 00000000..224f8cd9
--- /dev/null
+++ b/tests/test_data_communication.py
@@ -0,0 +1,85 @@
+from types import SimpleNamespace
+
+from pygwalker.communications.protocol import (
+ BatchPayloadQueryRequest,
+ BatchSqlQueryRequest,
+ PayloadQueryRequest,
+ SqlQueryRequest,
+)
+from pygwalker.services.data_communication import DataCommunicationService
+from pygwalker.services.global_var import GlobalVarManager
+
+
+def test_data_communication_queries_sql_and_payload():
+ calls = []
+ walker = SimpleNamespace(
+ data_parser=SimpleNamespace(
+ get_datas_by_sql=lambda sql: calls.append(("sql", sql)) or [{"total": 3}],
+ get_datas_by_payload=lambda payload: calls.append(("payload", payload)) or [{"city": "London"}],
+ )
+ )
+ service = DataCommunicationService(walker)
+
+ sql_response = service.get_datas(SqlQueryRequest(sql="SELECT 1"))
+ payload_response = service.get_datas_by_payload(
+ PayloadQueryRequest(payload={"workflow": [{"type": "view"}], "limit": 5})
+ )
+
+ assert sql_response == {"datas": [{"total": 3}]}
+ assert payload_response == {"datas": [{"city": "London"}]}
+ assert calls == [
+ ("sql", "SELECT 1"),
+ ("payload", {"workflow": [{"type": "view"}], "limit": 5}),
+ ]
+
+
+def test_data_communication_queries_batches():
+ calls = []
+ walker = SimpleNamespace(
+ data_parser=SimpleNamespace(
+ batch_get_datas_by_sql=lambda queries: calls.append(("sql", queries)) or [[{"total": 3}]],
+ batch_get_datas_by_payload=lambda payloads: calls.append(("payload", payloads)) or [[{"city": "London"}]],
+ )
+ )
+ service = DataCommunicationService(walker)
+
+ sql_response = service.batch_get_datas_by_sql(BatchSqlQueryRequest(queryList=["SELECT 1"]))
+ payload_response = service.batch_get_datas_by_payload(
+ BatchPayloadQueryRequest(queryList=[{"workflow": [{"type": "view"}], "offset": 2}])
+ )
+
+ assert sql_response == {"datas": [[{"total": 3}]]}
+ assert payload_response == {"datas": [[{"city": "London"}]]}
+ assert calls == [
+ ("sql", ["SELECT 1"]),
+ ("payload", [{"workflow": [{"type": "view"}], "offset": 2}]),
+ ]
+
+
+def test_data_communication_exports_dataframe_to_walker_and_global_state():
+ previous_exported_dataframe = GlobalVarManager.last_exported_dataframe
+ walker = SimpleNamespace(
+ _last_exported_dataframe=None,
+ data_parser=SimpleNamespace(
+ get_datas_by_sql=lambda _sql: [{"city": "Tokyo", "value": 2}],
+ get_datas_by_payload=lambda _payload: [{"city": "London", "value": 1}],
+ ),
+ )
+ service = DataCommunicationService(walker)
+
+ try:
+ sql_response = service.export_dataframe_by_sql(SqlQueryRequest(sql="SELECT *"))
+
+ assert sql_response == {}
+ assert walker._last_exported_dataframe.to_dict("records") == [{"city": "Tokyo", "value": 2}]
+ assert GlobalVarManager.last_exported_dataframe is walker._last_exported_dataframe
+
+ payload_response = service.export_dataframe_by_payload(
+ PayloadQueryRequest(payload={"workflow": [{"type": "view"}]})
+ )
+
+ assert payload_response == {}
+ assert walker._last_exported_dataframe.to_dict("records") == [{"city": "London", "value": 1}]
+ assert GlobalVarManager.last_exported_dataframe is walker._last_exported_dataframe
+ finally:
+ GlobalVarManager.last_exported_dataframe = previous_exported_dataframe
diff --git a/tests/test_data_parsers.py b/tests/test_data_parsers.py
index 9ba1eabd..a2306449 100644
--- a/tests/test_data_parsers.py
+++ b/tests/test_data_parsers.py
@@ -1,12 +1,15 @@
import os.path
+import subprocess
+import sys
from sqlalchemy import create_engine
import pandas as pd
import polars as pl
+import pyarrow as pa
import pytest
-from pygwalker.services.data_parsers import get_parser
-from pygwalker.data_parsers.database_parser import Connector, text
+from pygwalker.services.data_parsers import get_dataset_hash, get_parser
+from pygwalker.data_parsers.database_parser import Connector, DatabaseDataParser, text
from pygwalker.data_parsers.database_parser import _check_view_sql
from pygwalker.errors import ViewSqlSameColumnError
@@ -21,12 +24,12 @@
sql = "SELECT COUNT(1) total FROM pygwalker_mid_table"
sql_result = [{"total": 5}]
raw_fields_result = [
- {'fid': 'name', 'name': 'name', 'semanticType': 'nominal', 'analyticType': 'dimension'},
- {'fid': 'count', 'name': 'count', 'semanticType': 'quantitative', 'analyticType': 'dimension'},
- {'fid': 'date', 'name': 'date', 'semanticType': 'nominal', 'analyticType': 'dimension'}
+ {"fid": "name", "name": "name", "semanticType": "nominal", "analyticType": "dimension"},
+ {"fid": "count", "name": "count", "semanticType": "quantitative", "analyticType": "dimension"},
+ {"fid": "date", "name": "date", "semanticType": "nominal", "analyticType": "dimension"},
]
-to_records_result = [{'name': 'padnas', 'count': 3, 'date': '2022-01-01'}]
-to_records_no_kernel_result = [{'name': 'padnas', 'count': 3, 'date': '2022-01-01'}]
+to_records_result = [{"name": "padnas", "count": 3, "date": "2022-01-01"}]
+to_records_no_kernel_result = [{"name": "padnas", "count": 3, "date": "2022-01-01"}]
def test_data_parser_on_padnas():
@@ -49,8 +52,87 @@ def test_data_parser_on_polars():
assert dataset_parser.to_records(1) == to_records_no_kernel_result
+def test_data_parser_on_pyarrow_table():
+ table = pa.table({key: [row[key] for row in datas] for key in datas[0]})
+ dataset_parser = get_parser(table)
+
+ assert dataset_parser.dataset_type == "pyarrow_table"
+ assert dataset_parser.get_datas_by_sql(sql) == sql_result
+ assert dataset_parser.raw_fields == raw_fields_result
+ assert dataset_parser.to_records(1) == to_records_result
+ assert get_dataset_hash(table) == get_dataset_hash(table)
+
+
+@pytest.mark.parametrize(
+ "dataset",
+ [
+ pd.DataFrame({"city": pd.Series(dtype="object"), "value": pd.Series(dtype="int64")}),
+ pl.DataFrame({"city": pl.Series([], dtype=pl.String), "value": pl.Series([], dtype=pl.Int64)}),
+ pa.table({"city": pa.array([], type=pa.string()), "value": pa.array([], type=pa.int64())}),
+ ],
+)
+def test_data_parser_accepts_empty_tabular_inputs(dataset):
+ dataset_parser = get_parser(dataset)
+
+ assert dataset_parser.data_size == 0
+ assert dataset_parser.to_records() == []
+ assert dataset_parser.raw_fields == [
+ {"fid": "city", "name": "city", "semanticType": "nominal", "analyticType": "dimension"},
+ {"fid": "value", "name": "value", "semanticType": "quantitative", "analyticType": "dimension"},
+ ]
+
+
+def test_get_parser_reports_supported_inputs_for_unsupported_dataset():
+ with pytest.raises(TypeError) as exc_info:
+ get_parser(object())
+
+ message = str(exc_info.value)
+ assert "Unsupported dataset type: builtins.object" in message
+ assert "pandas.DataFrame" in message
+ assert "pyarrow.Table" in message
+ assert "pygwalker.data_parsers.database_parser.Connector" in message
+ assert "cloud dataset id string" in message
+
+
+@pytest.mark.parametrize(
+ "module_name",
+ [
+ "pygwalker.data_parsers.base",
+ "pygwalker.services.render_manager",
+ ],
+)
+def test_duckdb_import_failure_has_actionable_message(module_name):
+ repo_root = os.path.dirname(os.path.dirname(__file__))
+ code = f"""
+import builtins
+import importlib
+
+original_import = builtins.__import__
+
+def blocked_import(name, *args, **kwargs):
+ if name == "duckdb" or name.startswith("duckdb."):
+ raise ModuleNotFoundError("No module named 'duckdb'")
+ return original_import(name, *args, **kwargs)
+
+builtins.__import__ = blocked_import
+importlib.import_module({module_name!r})
+"""
+ result = subprocess.run(
+ [sys.executable, "-c", code],
+ cwd=repo_root,
+ check=False,
+ text=True,
+ capture_output=True,
+ )
+
+ assert result.returncode != 0
+ assert "PyGWalker requires duckdb for dataframe querying" in result.stderr
+ assert "pip install duckdb" in result.stderr
+
+
try:
from modin import pandas as mpd
+
def test_data_parser_on_modin():
df = mpd.DataFrame(datas)
dataset_parser = get_parser(df)
@@ -59,6 +141,25 @@ def test_data_parser_on_modin():
assert dataset_parser.to_records(1) == to_records_result
dataset_parser = get_parser(df)
assert dataset_parser.to_records(1) == to_records_no_kernel_result
+
+ def test_data_parser_accepts_empty_modin_input():
+ df = mpd.DataFrame({"city": [], "value": []})
+ dataset_parser = get_parser(df)
+
+ assert dataset_parser.data_size == 0
+ assert dataset_parser.to_records() == []
+ assert dataset_parser.raw_fields == [
+ {"fid": "city", "name": "city", "semanticType": "nominal", "analyticType": "dimension"},
+ {"fid": "value", "name": "value", "semanticType": "nominal", "analyticType": "dimension"},
+ ]
+
+ def test_data_parser_infers_modin_series_by_position_not_label():
+ df = mpd.DataFrame({"date": ["2022-01-01"]}, index=[10])
+ dataset_parser = get_parser(df, infer_string_to_date=True)
+
+ assert dataset_parser.raw_fields == [
+ {"fid": "date", "name": "date", "semanticType": "temporal", "analyticType": "dimension"},
+ ]
except ImportError:
pass
@@ -81,7 +182,7 @@ def test_check_view_sql():
def test_connector():
csv_file = os.path.join(os.path.dirname(__file__), "bike_sharing_dc.csv")
database_url = "duckdb:///:memory:"
- view_sql = f"SELECT 1"
+ view_sql = "SELECT 1"
data_count = 17379
connector = Connector(database_url, view_sql)
@@ -103,8 +204,26 @@ def test_connector():
with engine.connect() as conn:
conn.execute(text(f"CREATE TABLE test_datas AS SELECT * FROM read_csv_auto('{csv_file}')"))
connector = Connector.from_sqlalchemy_connection(conn, view_sql)
- result = connector.query_datas(f"SELECT COUNT(1) count FROM test_datas")
+ result = connector.query_datas("SELECT COUNT(1) count FROM test_datas")
assert result[0]["count"] == data_count
assert connector.dialect_name == "duckdb"
assert connector.view_sql == view_sql
assert connector.url == database_url
+
+
+def test_database_parser_get_datas_by_sql_queries_connector_view():
+ engine = create_engine("duckdb:///:memory:")
+ with engine.connect() as conn:
+ conn.execute(text("CREATE TABLE test_datas AS SELECT 1 AS id, 'London' AS city UNION ALL SELECT 2, 'Tokyo'"))
+ connector = Connector.from_sqlalchemy_connection(conn, "SELECT * FROM test_datas")
+ parser = DatabaseDataParser(connector, [], False, True, {})
+
+ assert parser.get_datas_by_sql("SELECT city FROM ___pygwalker_temp_view_name___ WHERE id = 2") == [
+ {"city": "Tokyo"}
+ ]
+ assert parser.batch_get_datas_by_sql(
+ [
+ "SELECT COUNT(1) AS total FROM ___pygwalker_temp_view_name___",
+ "SELECT city FROM ___pygwalker_temp_view_name___ WHERE id = 1",
+ ]
+ ) == [[{"total": 2}], [{"city": "London"}]]
diff --git a/tests/test_data_upload_communication.py b/tests/test_data_upload_communication.py
new file mode 100644
index 00000000..1fd51c7c
--- /dev/null
+++ b/tests/test_data_upload_communication.py
@@ -0,0 +1,23 @@
+from types import SimpleNamespace
+
+from pygwalker.services.data_upload_communication import DataUploadCommunicationService
+
+
+def test_data_upload_communication_requests_current_records():
+ upload_calls = []
+ walker = SimpleNamespace(
+ origin_data_source=[{"city": "London"}, {"city": "Tokyo"}],
+ data_source_id="data-source",
+ )
+ upload_tool = SimpleNamespace(run=lambda **kwargs: upload_calls.append(kwargs))
+
+ response = DataUploadCommunicationService(walker, upload_tool).request_data({})
+
+ assert response == {}
+ assert upload_calls == [
+ {
+ "records": [{"city": "London"}, {"city": "Tokyo"}],
+ "sample_data_count": 0,
+ "data_source_id": "data-source",
+ }
+ ]
diff --git a/tests/test_desktop_communication.py b/tests/test_desktop_communication.py
new file mode 100644
index 00000000..a10e0139
--- /dev/null
+++ b/tests/test_desktop_communication.py
@@ -0,0 +1,25 @@
+from types import SimpleNamespace
+
+from pygwalker.communications.protocol import OpenDesktopRequest
+from pygwalker.services.desktop_communication import DesktopCommunicationService
+
+
+def test_desktop_communication_imports_current_records():
+ import_calls = []
+ desktop_import = SimpleNamespace(import_to_desktop=lambda **kwargs: import_calls.append(kwargs))
+ walker = SimpleNamespace(
+ data_parser=SimpleNamespace(to_records=lambda: [{"city": "London"}, {"city": "Tokyo"}]),
+ )
+
+ response = DesktopCommunicationService(walker, desktop_import).open_in_desktop(
+ OpenDesktopRequest(spec=[{"name": "Chart"}], fields=[{"fid": "city"}])
+ )
+
+ assert response == {}
+ assert import_calls == [
+ {
+ "spec": [{"name": "Chart"}],
+ "fields": [{"fid": "city"}],
+ "records": [{"city": "London"}, {"city": "Tokyo"}],
+ }
+ ]
diff --git a/tests/test_desktop_import.py b/tests/test_desktop_import.py
new file mode 100644
index 00000000..2ab54ecd
--- /dev/null
+++ b/tests/test_desktop_import.py
@@ -0,0 +1,41 @@
+import base64
+import datetime as dt
+import json
+import urllib.parse
+import zlib
+
+from pygwalker.services.desktop_import import DesktopImportService
+
+
+def _decode_query_value(link: str, name: str):
+ parsed = urllib.parse.urlparse(link)
+ query = urllib.parse.parse_qs(parsed.query)
+ compressed = base64.b64decode(urllib.parse.unquote(query[name][0]))
+ return json.loads(zlib.decompress(compressed).decode())
+
+
+def test_desktop_import_service_builds_import_link():
+ service = DesktopImportService(open_link=lambda _link: None)
+
+ link = service.build_import_link(
+ spec=[{"name": "Chart"}],
+ fields=[{"fid": "city"}],
+ records=[{"city": "London", "date": dt.date(2024, 1, 1)}],
+ )
+
+ parsed = urllib.parse.urlparse(link)
+ assert parsed.scheme == "gw"
+ assert parsed.netloc == "import"
+ assert _decode_query_value(link, "spec") == [{"name": "Chart"}]
+ assert _decode_query_value(link, "fields") == [{"fid": "city"}]
+ assert _decode_query_value(link, "data") == [{"city": "London", "date": "2024-01-01"}]
+
+
+def test_desktop_import_service_opens_built_link():
+ links = []
+ service = DesktopImportService(open_link=links.append)
+
+ service.import_to_desktop(spec=[], fields=[], records=[])
+
+ assert len(links) == 1
+ assert urllib.parse.urlparse(links[0]).scheme == "gw"
diff --git a/tests/test_dsl_transform.py b/tests/test_dsl_transform.py
index 8739acf1..1bb15777 100644
--- a/tests/test_dsl_transform.py
+++ b/tests/test_dsl_transform.py
@@ -9,8 +9,7 @@
def _reset_runtime():
"""Reset the lazy-initialized JS runtime so tests are independent."""
- mod._dsl_to_workflow_js = None
- mod._vega_to_dsl_js = None
+ mod._close_js_runtime()
def test_dsl_to_workflow_returns_valid_workflow():
@@ -73,8 +72,29 @@ def counting_make(*args, **kwargs):
call_count[0] += 1
return original(*args, **kwargs)
- with mock.patch.object(mod, '_make_js_callable', side_effect=counting_make):
+ with mock.patch.object(mod, "_make_js_callable", side_effect=counting_make):
dsl_to_workflow({})
dsl_to_workflow({})
# _make_js_callable is called twice during init (once per UMD file), but only on first call
assert call_count[0] == 2
+
+
+def test_reset_runtime_closes_existing_contexts():
+ class Runtime:
+ def __init__(self):
+ self.closed = False
+
+ def close(self):
+ self.closed = True
+
+ dsl_runtime = Runtime()
+ vega_runtime = Runtime()
+ mod._dsl_to_workflow_js = dsl_runtime
+ mod._vega_to_dsl_js = vega_runtime
+
+ _reset_runtime()
+
+ assert dsl_runtime.closed is True
+ assert vega_runtime.closed is True
+ assert mod._dsl_to_workflow_js is None
+ assert mod._vega_to_dsl_js is None
diff --git a/tests/test_fname_encodings.py b/tests/test_fname_encodings.py
index 215509f1..132c1375 100644
--- a/tests/test_fname_encodings.py
+++ b/tests/test_fname_encodings.py
@@ -1,9 +1,4 @@
-from pygwalker.services.fname_encodings import (
- base36encode,
- fname_decode,
- fname_encode,
- base36decode
-)
+from pygwalker.services.fname_encodings import base36encode, fname_decode, fname_encode, base36decode
def test_base36_encode():
diff --git a/tests/test_frontend_assets.py b/tests/test_frontend_assets.py
new file mode 100644
index 00000000..972332d5
--- /dev/null
+++ b/tests/test_frontend_assets.py
@@ -0,0 +1,45 @@
+import pytest
+import json
+import ntpath
+from pathlib import Path
+
+from pygwalker.utils import frontend_assets
+
+
+def test_frontend_toolchain_targets_vite_6():
+ package_json = json.loads((Path(__file__).resolve().parents[1] / "app" / "package.json").read_text())
+
+ assert package_json["devDependencies"]["vite"].startswith("^6.")
+ assert package_json["devDependencies"]["@vitejs/plugin-react"].startswith("^4.")
+
+
+def test_read_frontend_asset_reads_from_templates_dist(monkeypatch, tmp_path):
+ asset_dir = tmp_path / "templates" / "dist"
+ asset_dir.mkdir(parents=True)
+ (asset_dir / "asset.js").write_text("console.log('ok');", encoding="utf-8")
+ monkeypatch.setattr(frontend_assets, "ROOT_DIR", str(tmp_path))
+
+ assert frontend_assets.read_frontend_asset("asset.js", encoding="utf-8") == "console.log('ok');"
+
+
+def test_read_frontend_asset_reports_compile_command(monkeypatch, tmp_path):
+ monkeypatch.setattr(frontend_assets, "ROOT_DIR", str(tmp_path))
+
+ with pytest.raises(RuntimeError) as exc_info:
+ frontend_assets.read_frontend_asset("missing.js")
+
+ message = str(exc_info.value)
+ assert "pygwalker/templates/dist/missing.js" in message
+ assert "./scripts/compile.sh" in message
+
+
+def test_read_frontend_asset_report_uses_stable_posix_path_with_windows_paths(monkeypatch, tmp_path):
+ monkeypatch.setattr(frontend_assets, "ROOT_DIR", str(tmp_path))
+ monkeypatch.setattr(frontend_assets.os, "path", ntpath)
+
+ with pytest.raises(RuntimeError) as exc_info:
+ frontend_assets.read_frontend_asset("missing.js")
+
+ message = str(exc_info.value)
+ assert "pygwalker/templates/dist/missing.js" in message
+ assert "pygwalker\\templates\\dist\\missing.js" not in message
diff --git a/tests/test_frontend_contracts.py b/tests/test_frontend_contracts.py
new file mode 100644
index 00000000..b8f47a9c
--- /dev/null
+++ b/tests/test_frontend_contracts.py
@@ -0,0 +1,342 @@
+import ast
+import re
+from pathlib import Path
+
+from pydantic import BaseModel
+
+from pygwalker.communications import protocol
+
+
+PYTHON_REQUEST_MODEL_TS_TYPES = {
+ "EmptyRequest": "ICommEmptyRequest",
+ "SqlQueryRequest": "ICommSqlQueryRequest",
+ "PayloadQueryRequest": "ICommPayloadQueryRequest",
+ "BatchSqlQueryRequest": "ICommBatchQueryRequest",
+ "BatchPayloadQueryRequest": "ICommBatchQueryRequest",
+ "UploadSpecToCloudRequest": "ICommUploadSpecToCloudRequest",
+ "SaveChartRequest": "ICommSaveChartRequest",
+ "UpdateSpecRequest": "ICommUpdateSpecRequest",
+ "AskSpecRequest": "ICommAskSpecRequest",
+ "ChatChartRequest": "ICommChatChartRequest",
+ "OpenDesktopRequest": "ICommOpenDesktopRequest",
+ "UploadCloudChartRequest": "ICommUploadCloudChartRequest",
+ "UploadCloudDashboardRequest": "ICommUploadCloudDashboardRequest",
+}
+
+PROTOCOL_MODEL_TS_INTERFACES = {
+ protocol.CommMessageRequest: "ICommEnvelope",
+ protocol.CommResponse: "ICommResponse",
+ protocol.EmptyRequest: "ICommEmptyRequest",
+ protocol.SqlQueryRequest: "ICommSqlQueryRequest",
+ protocol.PayloadQueryRequest: "ICommPayloadQueryRequest",
+ protocol.BatchSqlQueryRequest: "ICommBatchQueryRequest",
+ protocol.BatchPayloadQueryRequest: "ICommBatchQueryRequest",
+ protocol.UploadSpecToCloudRequest: "ICommUploadSpecToCloudRequest",
+ protocol.ChartImageRequest: "ICommChartImageRequest",
+ protocol.SaveChartRequest: "ICommSaveChartRequest",
+ protocol.UpdateSpecRequest: "ICommUpdateSpecRequest",
+ protocol.AskSpecRequest: "ICommAskSpecRequest",
+ protocol.ChatChartRequest: "ICommChatChartRequest",
+ protocol.OpenDesktopRequest: "ICommOpenDesktopRequest",
+ protocol.UploadCloudChartRequest: "ICommUploadCloudChartRequest",
+ protocol.UploadCloudDashboardRequest: "ICommUploadCloudDashboardRequest",
+ protocol.EmptyResponse: "ICommEmptyResponse",
+ protocol.LatestVisSpecResponse: "ICommLatestVisSpecResponse",
+ protocol.DataRowsResponse: "ICommDataRowsResponse",
+ protocol.BatchDataRowsResponse: "ICommBatchDataRowsResponse",
+ protocol.UploadSpecToCloudResponse: "ICommUploadSpecToCloudResponse",
+ protocol.CloudCallbackResponse: "ICommCloudCallbackResponse",
+ protocol.UploadCloudChartResponse: "ICommUploadCloudChartResponse",
+ protocol.UploadCloudDashboardResponse: "ICommUploadCloudDashboardResponse",
+}
+
+
+def _props_builder_keys(repo_root: Path) -> set[str]:
+ source = (repo_root / "pygwalker/services/props_builder.py").read_text(encoding="utf-8")
+ module = ast.parse(source)
+
+ for class_node in module.body:
+ if not isinstance(class_node, ast.ClassDef) or class_node.name != "PropsBuilder":
+ continue
+ for method_node in class_node.body:
+ if not isinstance(method_node, ast.FunctionDef) or method_node.name != "build":
+ continue
+ for node in ast.walk(method_node):
+ if isinstance(node, ast.Return) and isinstance(node.value, ast.Dict):
+ return {
+ key.value
+ for key in node.value.keys
+ if isinstance(key, ast.Constant) and isinstance(key.value, str)
+ }
+ raise AssertionError("Could not find PropsBuilder.build return keys")
+
+
+def _app_props_keys(repo_root: Path) -> set[str]:
+ source = (repo_root / "app/src/interfaces/index.ts").read_text(encoding="utf-8")
+ match = re.search(r"export interface IAppProps \{([\s\S]*?)\n\}", source)
+ if match is None:
+ raise AssertionError("Could not find IAppProps interface")
+ return set(re.findall(r"^\s*([A-Za-z_][A-Za-z0-9_]*)\??:", match.group(1), re.MULTILINE))
+
+
+def _comm_handler_endpoints(repo_root: Path) -> set[str]:
+ source = (repo_root / "pygwalker/services/comm_handler.py").read_text(encoding="utf-8")
+ module = ast.parse(source)
+
+ for class_node in module.body:
+ if not isinstance(class_node, ast.ClassDef) or class_node.name != "CommHandler":
+ continue
+ for method_node in class_node.body:
+ if not isinstance(method_node, ast.FunctionDef) or method_node.name != "register":
+ continue
+ endpoints = set()
+ for node in ast.walk(method_node):
+ if not isinstance(node, ast.Call):
+ continue
+ if not isinstance(node.func, ast.Attribute):
+ continue
+ if node.func.attr not in {"register", "_register_request"}:
+ continue
+ if node.args and isinstance(node.args[0], ast.Constant) and isinstance(node.args[0].value, str):
+ endpoints.add(node.args[0].value)
+ return endpoints
+ raise AssertionError("Could not find CommHandler.register endpoints")
+
+
+def _comm_handler_request_models(repo_root: Path) -> dict[str, str]:
+ source = (repo_root / "pygwalker/services/comm_handler.py").read_text(encoding="utf-8")
+ module = ast.parse(source)
+
+ for class_node in module.body:
+ if not isinstance(class_node, ast.ClassDef) or class_node.name != "CommHandler":
+ continue
+ for method_node in class_node.body:
+ if not isinstance(method_node, ast.FunctionDef) or method_node.name != "register":
+ continue
+ endpoints = {}
+ for node in ast.walk(method_node):
+ if not isinstance(node, ast.Call):
+ continue
+ if not isinstance(node.func, ast.Attribute) or node.func.attr != "_register_request":
+ continue
+ if len(node.args) < 2:
+ continue
+ endpoint_node, model_node = node.args[0], node.args[1]
+ if isinstance(endpoint_node, ast.Constant) and isinstance(endpoint_node.value, str):
+ if isinstance(model_node, ast.Name):
+ endpoints[endpoint_node.value] = model_node.id
+ return endpoints
+ raise AssertionError("Could not find CommHandler.register request models")
+
+
+def _comm_handler_register_calls(repo_root: Path) -> list[ast.Call]:
+ source = (repo_root / "pygwalker/services/comm_handler.py").read_text(encoding="utf-8")
+ module = ast.parse(source)
+
+ for class_node in module.body:
+ if not isinstance(class_node, ast.ClassDef) or class_node.name != "CommHandler":
+ continue
+ for method_node in class_node.body:
+ if isinstance(method_node, ast.FunctionDef) and method_node.name == "register":
+ return [node for node in ast.walk(method_node) if isinstance(node, ast.Call)]
+ raise AssertionError("Could not find CommHandler.register")
+
+
+def _typescript_interface_keys(repo_root: Path, interface_name: str) -> set[str]:
+ source = (repo_root / "app/src/interfaces/index.ts").read_text(encoding="utf-8")
+ match = re.search(rf"export interface {interface_name} \{{([\s\S]*?)\n\}}", source)
+ if match is None:
+ raise AssertionError(f"Could not find {interface_name} interface")
+ return set(re.findall(r"^\s*([A-Za-z_][A-Za-z0-9_]*)\??:", match.group(1), re.MULTILINE))
+
+
+def _typescript_shape_keys(repo_root: Path, type_name: str) -> set[str]:
+ source = (repo_root / "app/src/interfaces/index.ts").read_text(encoding="utf-8")
+ empty_interface_pattern = rf"export interface {type_name}(?:<[^>]+>)? \{{\}}"
+ if re.search(empty_interface_pattern, source) is not None:
+ return set()
+
+ interface_match = re.search(rf"export interface {type_name}(?:<[^>]+>)? \{{([\s\S]*?)\n\}}", source)
+ if interface_match is not None:
+ return set(re.findall(r"^\s*([A-Za-z_][A-Za-z0-9_]*)\??:", interface_match.group(1), re.MULTILINE))
+
+ empty_alias_pattern = rf"export type {type_name} = Record;"
+ if re.search(empty_alias_pattern, source) is not None:
+ return set()
+
+ raise AssertionError(f"Could not find TypeScript shape for {type_name}")
+
+
+def _typescript_map_values(repo_root: Path, interface_name: str) -> dict[str, str]:
+ source = (repo_root / "app/src/interfaces/index.ts").read_text(encoding="utf-8")
+ match = re.search(rf"export interface {interface_name} \{{([\s\S]*?)\n\}}", source)
+ if match is None:
+ raise AssertionError(f"Could not find {interface_name} interface")
+ return dict(
+ re.findall(
+ r"^\s*([A-Za-z_][A-Za-z0-9_]*):\s*([^;]+);",
+ match.group(1),
+ re.MULTILINE,
+ )
+ )
+
+
+def _pydantic_alias_keys(model_cls: type[BaseModel]) -> set[str]:
+ fields = getattr(model_cls, "model_fields", None)
+ if fields is not None:
+ return {field.alias or name for name, field in fields.items()}
+ return {field.alias or name for name, field in model_cls.__fields__.items()}
+
+
+def test_frontend_props_interface_covers_python_props_builder():
+ repo_root = Path(__file__).resolve().parents[1]
+
+ assert _props_builder_keys(repo_root) - _app_props_keys(repo_root) == set()
+
+
+def test_python_props_builder_stringifies_frontend_id():
+ repo_root = Path(__file__).resolve().parents[1]
+ props_builder_source = (repo_root / "pygwalker/services/props_builder.py").read_text(encoding="utf-8")
+
+ assert '"id": str(self.walker.gid)' in props_builder_source
+
+
+def test_frontend_comm_maps_cover_python_comm_handler_endpoints():
+ repo_root = Path(__file__).resolve().parents[1]
+ endpoints = _comm_handler_endpoints(repo_root)
+
+ assert _typescript_interface_keys(repo_root, "ICommRequestMap") == endpoints
+ assert _typescript_interface_keys(repo_root, "ICommResponseMap") == endpoints
+
+
+def test_frontend_comm_request_map_uses_types_matching_python_request_models():
+ repo_root = Path(__file__).resolve().parents[1]
+ request_models = _comm_handler_request_models(repo_root)
+ request_map = _typescript_map_values(repo_root, "ICommRequestMap")
+
+ assert {
+ endpoint: PYTHON_REQUEST_MODEL_TS_TYPES[model_name] for endpoint, model_name in request_models.items()
+ } == request_map
+
+
+def test_frontend_protocol_interfaces_match_pydantic_aliases():
+ repo_root = Path(__file__).resolve().parents[1]
+
+ for model_cls, interface_name in PROTOCOL_MODEL_TS_INTERFACES.items():
+ assert _typescript_shape_keys(repo_root, interface_name) == _pydantic_alias_keys(model_cls)
+
+
+def test_comm_handler_register_uses_typed_request_models():
+ repo_root = Path(__file__).resolve().parents[1]
+ raw_register_calls = []
+ untyped_request_calls = []
+
+ for call in _comm_handler_register_calls(repo_root):
+ if not isinstance(call.func, ast.Attribute):
+ continue
+ if call.func.attr == "register":
+ raw_register_calls.append(call.lineno)
+ if call.func.attr != "_register_request":
+ continue
+ if len(call.args) < 2 or not isinstance(call.args[1], ast.Name) or not call.args[1].id.endswith("Request"):
+ untyped_request_calls.append(call.lineno)
+
+ assert raw_register_calls == []
+ assert untyped_request_calls == []
+
+
+def test_frontend_dev_typescript_source_maps_are_enabled():
+ repo_root = Path(__file__).resolve().parents[1]
+ tsconfig = (repo_root / "app/tsconfig.json").read_text(encoding="utf-8")
+
+ assert '"sourceMap": true' in tsconfig
+ assert '"sourceMap": false' not in tsconfig
+
+
+def test_frontend_tracker_loads_segment_only_after_events_opt_in():
+ repo_root = Path(__file__).resolve().parents[1]
+ tracker_source = (repo_root / "app/src/utils/tracker.ts").read_text(encoding="utf-8")
+ app_source = (repo_root / "app/src/index.tsx").read_text(encoding="utf-8")
+
+ assert "@segment/analytics-next" not in tracker_source
+ assert "cdn.segment.com/analytics.js/v1/" in tracker_source
+ assert 'tracker.setOpen(userConfig.privacy === "events")' in app_source
+
+
+def test_frontend_http_integrations_initialize_communication():
+ repo_root = Path(__file__).resolve().parents[1]
+ app_source = (repo_root / "app/src/index.tsx").read_text(encoding="utf-8")
+
+ for env in ("streamlit", "gradio", "web_server"):
+ pattern = rf'case "{env}":[\s\S]*?preRender = initOnHttpCommunication;'
+ assert re.search(pattern, app_source) is not None
+
+ assert 'case "reflex":' not in app_source
+
+
+def test_streamlit_entrypoint_refreshes_props_on_reruns():
+ repo_root = Path(__file__).resolve().parents[1]
+ app_source = (repo_root / "app/src/index.tsx").read_text(encoding="utf-8")
+ streamlit_app_match = re.search(
+ r"function SteamlitGWalkerApp[\s\S]*?const StreamlitGWalker =",
+ app_source,
+ )
+
+ assert streamlit_app_match is not None
+ streamlit_app_source = streamlit_app_match.group(0)
+ assert "propsRef" not in streamlit_app_source
+ assert "formatAppProps(streamlitProps.args as IAppProps)" in streamlit_app_source
+ assert "[streamlitProps.args]" in streamlit_app_source
+
+
+def test_frontend_component_state_refreshes_when_props_change():
+ repo_root = Path(__file__).resolve().parents[1]
+ app_source = (repo_root / "app/src/index.tsx").read_text(encoding="utf-8")
+
+ assert "setDataSource(props.dataSource);" in app_source
+ assert "setVisSpec(props.visSpec);" in app_source
+
+
+def test_frontend_pure_renderer_uses_local_data_without_kernel_computation():
+ repo_root = Path(__file__).resolve().parents[1]
+ app_source = (repo_root / "app/src/index.tsx").read_text(encoding="utf-8")
+
+ pure_renderer_match = re.search(
+ r"const PureRednererApp:[\s\S]*?const initOnJupyter",
+ app_source,
+ )
+ assert pure_renderer_match is not None
+ pure_renderer_source = pure_renderer_match.group(0)
+
+ assert "props.useKernelCalc ?" in pure_renderer_source
+ assert "type='remote'" in pure_renderer_source
+ assert "computation={computationCallback!}" in pure_renderer_source
+ assert "rawData={props.dataSource}" in pure_renderer_source
+
+
+def test_frontend_modals_are_lazy_loaded_from_entrypoint():
+ repo_root = Path(__file__).resolve().parents[1]
+ app_source = (repo_root / "app/src/index.tsx").read_text(encoding="utf-8")
+
+ modal_paths = [
+ "./components/initModal",
+ "./components/uploadSpecModal",
+ "./components/uploadChartModal",
+ "./components/codeExportModal",
+ ]
+ for modal_path in modal_paths:
+ static_import_pattern = rf"import\s+[^;\n]+?\s+from\s+[\"']{re.escape(modal_path)}[\"']"
+
+ assert re.search(static_import_pattern, app_source) is None
+ assert f'import "{modal_path}"' not in app_source
+ assert f"import '{modal_path}'" not in app_source
+ assert f'import("{modal_path}")' in app_source
+
+
+def test_frontend_save_payload_strips_graphic_walker_export_fields():
+ repo_root = Path(__file__).resolve().parents[1]
+ save_source = (repo_root / "app/src/utils/save.ts").read_text(encoding="utf-8")
+
+ assert "Promise" in save_source
+ assert "...chartData" not in save_source
+ assert "canvas: () => null" not in save_source
diff --git a/tests/test_integration_apis.py b/tests/test_integration_apis.py
new file mode 100644
index 00000000..f74b099a
--- /dev/null
+++ b/tests/test_integration_apis.py
@@ -0,0 +1,852 @@
+import importlib
+import json
+import sys
+from types import ModuleType
+from types import SimpleNamespace
+
+import pandas as pd
+import pytest
+
+
+_OPTIONAL_INTEGRATION_MODULES = [
+ "anywidget",
+ "marimo",
+ "reflex",
+ "streamlit",
+ "streamlit.components",
+ "streamlit.components.v1",
+ "streamlit.components.v1.components",
+ "pygwalker.api.anywidget",
+ "pygwalker.api.marimo",
+ "pygwalker.api.reflex",
+ "pygwalker.api.streamlit",
+ "pygwalker.communications.anywidget_comm",
+ "pygwalker.communications.streamlit_comm",
+ "pygwalker.services.anywidget_widget",
+ "pygwalker.services.streamlit_components",
+]
+
+_OPTIONAL_PARENT_ATTRS = [
+ ("pygwalker.api", "anywidget"),
+ ("pygwalker.api", "marimo"),
+ ("pygwalker.api", "reflex"),
+ ("pygwalker.api", "streamlit"),
+ ("pygwalker.communications", "anywidget_comm"),
+ ("pygwalker.communications", "streamlit_comm"),
+ ("pygwalker.services", "anywidget_widget"),
+ ("pygwalker.services", "streamlit_components"),
+ ("streamlit", "components"),
+ ("streamlit.components", "v1"),
+]
+
+
+@pytest.fixture(autouse=True)
+def restore_optional_integration_modules():
+ missing = object()
+ snapshot = {name: sys.modules.get(name, missing) for name in _OPTIONAL_INTEGRATION_MODULES}
+ parent_snapshot = {}
+ for module_name, attr_name in _OPTIONAL_PARENT_ATTRS:
+ module = sys.modules.get(module_name)
+ if module is None:
+ parent_snapshot[(module_name, attr_name)] = (missing, missing)
+ else:
+ parent_snapshot[(module_name, attr_name)] = (module, getattr(module, attr_name, missing))
+ if hasattr(module, attr_name):
+ delattr(module, attr_name)
+
+ for name in _OPTIONAL_INTEGRATION_MODULES:
+ sys.modules.pop(name, None)
+
+ yield
+ for name in _OPTIONAL_INTEGRATION_MODULES:
+ sys.modules.pop(name, None)
+ for name, module in snapshot.items():
+ if module is not missing:
+ sys.modules[name] = module
+
+ for (module_name, attr_name), (module, value) in parent_snapshot.items():
+ if module is missing:
+ continue
+ if value is missing:
+ if hasattr(module, attr_name):
+ delattr(module, attr_name)
+ else:
+ setattr(module, attr_name, value)
+
+
+class FakeWalker:
+ instances = []
+
+ def __init__(self, **kwargs):
+ self.kwargs = kwargs
+ self.gid = kwargs["gid"] or "generated"
+ self.init_callback_calls = []
+ self.props_calls = []
+ self.render_calls = []
+ self.kernel_computation = kwargs["kernel_computation"]
+ self.origin_data_source = [{"city": "London"}]
+ FakeWalker.instances.append(self)
+
+ def _get_props(self, env="", data_source=None, need_load_datas=False):
+ self.props_calls.append((env, data_source, need_load_datas))
+ return {
+ "id": self.gid,
+ "env": env,
+ "dataSource": data_source,
+ }
+
+ def _init_callback(self, comm, preview_tool=None):
+ self.init_callback_calls.append((comm, preview_tool))
+
+ def _get_render_iframe(self, props, return_iframe=True, iframe_width=None, iframe_height=None):
+ self.render_calls.append((props, return_iframe, iframe_width, iframe_height))
+ return json.dumps(props, sort_keys=True)
+
+
+def _reset_fake_walker():
+ FakeWalker.instances = []
+
+
+def test_gradio_api_builds_walker_and_registers_comm(monkeypatch, tmp_path):
+ from pygwalker.api import gradio
+
+ _reset_fake_walker()
+ monkeypatch.setattr(gradio, "PygWalker", FakeWalker)
+ comms = []
+
+ class FakeGradioCommunication:
+ def __init__(self, gid):
+ self.gid = gid
+ comms.append(self)
+
+ monkeypatch.setattr(gradio, "GradioCommunication", FakeGradioCommunication)
+
+ html = gradio.get_html_on_gradio(
+ pd.DataFrame([{"city": "London"}]),
+ gid="gradio",
+ spec_path=str(tmp_path / "gradio_spec.json"),
+ spec_io_mode="rw",
+ computation="browser",
+ default_tab="data",
+ )
+
+ walker = FakeWalker.instances[0]
+ assert walker.kwargs["gid"] == "gradio"
+ assert walker.kwargs["spec"] == str(tmp_path / "gradio_spec.json")
+ assert walker.kwargs["use_save_tool"] is True
+ assert walker.kwargs["kernel_computation"] is False
+ assert walker.kwargs["default_tab"] == "data"
+ assert walker.props_calls == [("gradio", None, False)]
+ assert walker.init_callback_calls == [(comms[0], None)]
+ assert json.loads(html)["communicationUrl"] == gradio.BASE_URL_PATH
+
+
+def test_reflex_api_returns_html_component(monkeypatch):
+ fake_reflex = SimpleNamespace(Component=object, html=lambda value: {"html": value})
+ monkeypatch.setitem(sys.modules, "reflex", fake_reflex)
+ reflex = importlib.reload(importlib.import_module("pygwalker.api.reflex"))
+
+ _reset_fake_walker()
+ monkeypatch.setattr(reflex, "PygWalker", FakeWalker)
+ comms = []
+
+ class FakeReflexCommunication:
+ def __init__(self, gid):
+ self.gid = gid
+ comms.append(self)
+
+ monkeypatch.setattr(reflex, "ReflexCommunication", FakeReflexCommunication)
+
+ component = reflex.get_component(
+ pd.DataFrame([{"city": "London"}]),
+ gid="reflex",
+ spec_io_mode="r",
+ computation="browser",
+ )
+
+ walker = FakeWalker.instances[0]
+ assert walker.kwargs["gid"] == "reflex"
+ assert walker.kwargs["use_save_tool"] is False
+ assert walker.kwargs["kernel_computation"] is False
+ assert walker.props_calls == [("reflex", None, False)]
+ assert walker.init_callback_calls == [(comms[0], None)]
+ assert json.loads(component["html"])["communicationUrl"] == reflex.BASE_URL_PATH
+
+
+@pytest.mark.parametrize("computation", ["kernel", "cloud"])
+def test_reflex_api_rejects_live_computation(monkeypatch, computation):
+ fake_reflex = SimpleNamespace(Component=object, html=lambda value: {"html": value})
+ monkeypatch.setitem(sys.modules, "reflex", fake_reflex)
+ reflex = importlib.reload(importlib.import_module("pygwalker.api.reflex"))
+
+ with pytest.raises(ValueError, match="Reflex integration does not support kernel or cloud computation"):
+ reflex.get_component(
+ pd.DataFrame([{"city": "London"}]),
+ computation=computation,
+ )
+
+
+def test_anywidget_api_builds_widget_props_and_registers_comm(monkeypatch):
+ _install_anywidget_stubs(monkeypatch)
+ anywidget_api = importlib.reload(importlib.import_module("pygwalker.api.anywidget"))
+
+ _reset_fake_walker()
+ monkeypatch.setattr(anywidget_api, "PygWalker", FakeWalker)
+ comms = []
+
+ class FakeAnywidgetCommunication:
+ def __init__(self, gid):
+ self.gid = gid
+ self.widgets = []
+ comms.append(self)
+
+ def register_widget(self, widget):
+ self.widgets.append(widget)
+
+ monkeypatch.setattr(anywidget_api, "AnywidgetCommunication", FakeAnywidgetCommunication)
+
+ widget = anywidget_api.walk(
+ pd.DataFrame([{"city": "London"}]),
+ gid="anywidget",
+ show_cloud_tool=True,
+ computation="browser",
+ )
+
+ walker = FakeWalker.instances[0]
+ assert walker.kwargs["gid"] == "anywidget"
+ assert walker.kwargs["kernel_computation"] is False
+ assert walker.kwargs["use_save_tool"] is True
+ assert walker.kwargs["is_export_dataframe"] is True
+ assert walker.props_calls == [("anywidget", [{"city": "London"}], False)]
+ assert json.loads(widget.props)["env"] == "anywidget"
+ assert comms[0].widgets == [widget]
+ assert walker.init_callback_calls == [(comms[0], None)]
+
+
+def test_anywidget_api_accepts_public_walker_object(monkeypatch):
+ _install_anywidget_stubs(monkeypatch)
+ anywidget_api = importlib.reload(importlib.import_module("pygwalker.api.anywidget"))
+ walker_api = importlib.import_module("pygwalker.api.walker")
+
+ _reset_fake_walker()
+ monkeypatch.setattr(walker_api, "PygWalker", FakeWalker)
+ comms = []
+
+ class FakeAnywidgetCommunication:
+ def __init__(self, gid):
+ self.gid = gid
+ self.widgets = []
+ comms.append(self)
+
+ def register_widget(self, widget):
+ self.widgets.append(widget)
+
+ monkeypatch.setattr(anywidget_api, "AnywidgetCommunication", FakeAnywidgetCommunication)
+
+ public_walker = walker_api.Walker(
+ pd.DataFrame([{"city": "London"}]),
+ gid="anywidget-core",
+ computation="browser",
+ )
+ widget = anywidget_api.walk(public_walker)
+
+ assert len(FakeWalker.instances) == 1
+ assert public_walker.core.use_preview is False
+ assert public_walker.core.props_calls == [("anywidget", [{"city": "London"}], False)]
+ assert json.loads(widget.props)["env"] == "anywidget"
+ assert comms[0].widgets == [widget]
+ assert public_walker.core.init_callback_calls == [(comms[0], None)]
+
+
+def test_anywidget_api_rejects_rebuilding_public_walker_object(monkeypatch, tmp_path):
+ _install_anywidget_stubs(monkeypatch)
+ anywidget_api = importlib.reload(importlib.import_module("pygwalker.api.anywidget"))
+ walker_api = importlib.import_module("pygwalker.api.walker")
+
+ _reset_fake_walker()
+ monkeypatch.setattr(walker_api, "PygWalker", FakeWalker)
+ public_walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), computation="browser")
+
+ with pytest.raises(ValueError, match="cannot apply construction parameters: spec_path"):
+ anywidget_api.walk(public_walker, spec_path=str(tmp_path / "other.json"))
+
+
+def test_anywidget_api_rejects_show_cloud_tool_false_alias_for_public_walker(monkeypatch):
+ _install_anywidget_stubs(monkeypatch)
+ anywidget_api = importlib.reload(importlib.import_module("pygwalker.api.anywidget"))
+ walker_api = importlib.import_module("pygwalker.api.walker")
+
+ _reset_fake_walker()
+ monkeypatch.setattr(walker_api, "PygWalker", FakeWalker)
+ public_walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), computation="browser")
+
+ with pytest.raises(ValueError, match="cannot apply construction parameters: show_cloud_tool"):
+ anywidget_api.walk(public_walker, show_cloud_tool=0)
+
+
+def test_anywidget_widget_service_registers_comm(monkeypatch):
+ _install_anywidget_stubs(monkeypatch)
+ anywidget_widget = importlib.reload(importlib.import_module("pygwalker.services.anywidget_widget"))
+
+ class FakeWalkerForWidget:
+ gid = "widget-core"
+
+ def __init__(self):
+ self.props_calls = []
+ self.init_callback_calls = []
+
+ def _get_props(self, env, data_source):
+ self.props_calls.append((env, data_source))
+ return {"id": self.gid, "env": env, "dataSource": data_source}
+
+ def _init_callback(self, comm):
+ self.init_callback_calls.append(comm)
+
+ comms = []
+
+ class FakeAnywidgetCommunication:
+ def __init__(self, gid):
+ self.gid = gid
+ self.widgets = []
+ comms.append(self)
+
+ def register_widget(self, widget):
+ self.widgets.append(widget)
+
+ walker = FakeWalkerForWidget()
+ widget = anywidget_widget.create_anywidget_for_walker(
+ walker,
+ env="anywidget",
+ data_source=[],
+ communication_cls=FakeAnywidgetCommunication,
+ )
+
+ assert json.loads(widget.props) == {"id": "widget-core", "env": "anywidget", "dataSource": []}
+ assert walker.props_calls == [("anywidget", [])]
+ assert comms[0].gid == "widget-core"
+ assert comms[0].widgets == [widget]
+ assert walker.init_callback_calls == [comms[0]]
+
+
+def test_anywidget_widget_service_serializes_dataframe_values(monkeypatch):
+ _install_anywidget_stubs(monkeypatch)
+ anywidget_widget = importlib.reload(importlib.import_module("pygwalker.services.anywidget_widget"))
+
+ class FakeWalkerForWidget:
+ gid = "widget-core"
+
+ def _get_props(self, env, data_source):
+ return {
+ "id": self.gid,
+ "env": env,
+ "dataSource": [
+ {"started_at": pd.Timestamp("2024-01-01T00:00:00Z")},
+ {"started_at": pd.Timestamp("2024-01-02").date()},
+ ],
+ }
+
+ def _init_callback(self, comm):
+ self.comm = comm
+
+ class FakeAnywidgetCommunication:
+ def __init__(self, gid):
+ self.gid = gid
+
+ def register_widget(self, widget):
+ self.widget = widget
+
+ widget = anywidget_widget.create_anywidget_for_walker(
+ FakeWalkerForWidget(),
+ env="anywidget",
+ data_source=[],
+ communication_cls=FakeAnywidgetCommunication,
+ )
+
+ props = json.loads(widget.props)
+ assert props["dataSource"] == [
+ {"started_at": 1704067200000},
+ {"started_at": "2024-01-02"},
+ ]
+
+
+def test_marimo_api_wraps_anywidget(monkeypatch):
+ wrapped_widgets = []
+ _install_anywidget_stubs(monkeypatch)
+ monkeypatch.setitem(
+ sys.modules,
+ "marimo",
+ SimpleNamespace(
+ ui=SimpleNamespace(anywidget=lambda widget: wrapped_widgets.append(widget) or {"wrapped": widget})
+ ),
+ )
+ marimo_api = importlib.reload(importlib.import_module("pygwalker.api.marimo"))
+
+ _reset_fake_walker()
+ monkeypatch.setattr(marimo_api, "PygWalker", FakeWalker)
+ comms = []
+
+ class FakeAnywidgetCommunication:
+ def __init__(self, gid):
+ self.gid = gid
+ self.widgets = []
+ comms.append(self)
+
+ def register_widget(self, widget):
+ self.widgets.append(widget)
+
+ monkeypatch.setattr(marimo_api, "AnywidgetCommunication", FakeAnywidgetCommunication)
+
+ result = marimo_api.walk(pd.DataFrame([{"city": "London"}]), gid="marimo")
+
+ walker = FakeWalker.instances[0]
+ assert walker.props_calls == [("marimo", [], False)]
+ assert json.loads(wrapped_widgets[0].props)["env"] == "marimo"
+ assert comms[0].widgets == wrapped_widgets
+ assert walker.init_callback_calls == [(comms[0], None)]
+ assert result == {"wrapped": wrapped_widgets[0]}
+
+
+def test_marimo_api_accepts_public_walker_object(monkeypatch):
+ wrapped_widgets = []
+ _install_anywidget_stubs(monkeypatch)
+ monkeypatch.setitem(
+ sys.modules,
+ "marimo",
+ SimpleNamespace(
+ ui=SimpleNamespace(anywidget=lambda widget: wrapped_widgets.append(widget) or {"wrapped": widget})
+ ),
+ )
+ marimo_api = importlib.reload(importlib.import_module("pygwalker.api.marimo"))
+ walker_api = importlib.import_module("pygwalker.api.walker")
+
+ _reset_fake_walker()
+ monkeypatch.setattr(walker_api, "PygWalker", FakeWalker)
+ comms = []
+
+ class FakeAnywidgetCommunication:
+ def __init__(self, gid):
+ self.gid = gid
+ self.widgets = []
+ comms.append(self)
+
+ def register_widget(self, widget):
+ self.widgets.append(widget)
+
+ monkeypatch.setattr(marimo_api, "AnywidgetCommunication", FakeAnywidgetCommunication)
+
+ public_walker = walker_api.Walker(
+ pd.DataFrame([{"city": "London"}]),
+ gid="marimo-core",
+ computation="browser",
+ )
+ result = marimo_api.walk(public_walker)
+
+ assert len(FakeWalker.instances) == 1
+ assert public_walker.core.use_preview is False
+ assert public_walker.core.props_calls == [("marimo", [{"city": "London"}], False)]
+ assert json.loads(wrapped_widgets[0].props)["env"] == "marimo"
+ assert comms[0].widgets == wrapped_widgets
+ assert public_walker.core.init_callback_calls == [(comms[0], None)]
+ assert result == {"wrapped": wrapped_widgets[0]}
+
+
+def test_marimo_api_rejects_rebuilding_public_walker_object(monkeypatch, tmp_path):
+ _install_anywidget_stubs(monkeypatch)
+ monkeypatch.setitem(sys.modules, "marimo", SimpleNamespace(ui=SimpleNamespace(anywidget=lambda widget: widget)))
+ marimo_api = importlib.reload(importlib.import_module("pygwalker.api.marimo"))
+ walker_api = importlib.import_module("pygwalker.api.walker")
+
+ _reset_fake_walker()
+ monkeypatch.setattr(walker_api, "PygWalker", FakeWalker)
+ public_walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), computation="browser")
+
+ with pytest.raises(ValueError, match="cannot apply construction parameters: spec_path"):
+ marimo_api.walk(public_walker, spec_path=str(tmp_path / "other.json"))
+
+
+def test_streamlit_html_builds_renderer_and_component_html(monkeypatch, tmp_path):
+ _install_streamlit_stubs(monkeypatch)
+ streamlit = importlib.reload(importlib.import_module("pygwalker.api.streamlit"))
+
+ _reset_fake_walker()
+ monkeypatch.setattr(streamlit, "PygWalker", FakeWalker)
+ monkeypatch.setattr(streamlit, "init_streamlit_comm", lambda: None)
+ monkeypatch.setattr(streamlit, "get_dataset_hash", lambda _dataset: "dataset-hash")
+ monkeypatch.setattr(streamlit, "StreamlitCommunication", lambda gid: {"gid": gid})
+
+ html = streamlit.get_streamlit_html(
+ pd.DataFrame([{"city": "London"}]),
+ gid=None,
+ spec_path=str(tmp_path / "streamlit_spec.json"),
+ spec_io_mode="rw",
+ computation="cloud",
+ mode="table",
+ default_tab="data",
+ )
+
+ walker = FakeWalker.instances[0]
+ assert walker.kwargs["gid"] == "dataset-hash"
+ assert walker.kwargs["spec"] == str(tmp_path / "streamlit_spec.json")
+ assert walker.kwargs["use_save_tool"] is True
+ assert walker.kwargs["kernel_computation"] is False
+ assert walker.kwargs["cloud_computation"] is True
+ assert walker.kwargs["default_tab"] == "data"
+ assert walker.init_callback_calls == [({"gid": "dataset-hash"}, None)]
+ rendered_props = json.loads(html)
+ assert rendered_props["env"] == "streamlit"
+ assert rendered_props["communicationUrl"] == streamlit.BASE_URL_PATH
+ assert rendered_props["gwMode"] == "table"
+
+
+def test_streamlit_renderer_accepts_public_walker_object(monkeypatch):
+ _install_streamlit_stubs(monkeypatch)
+ streamlit = importlib.reload(importlib.import_module("pygwalker.api.streamlit"))
+ walker_api = importlib.import_module("pygwalker.api.walker")
+
+ _reset_fake_walker()
+ monkeypatch.setattr(walker_api, "PygWalker", FakeWalker)
+ monkeypatch.setattr(streamlit, "init_streamlit_comm", lambda: None)
+ monkeypatch.setattr(streamlit, "StreamlitCommunication", lambda gid: {"gid": gid})
+
+ public_walker = walker_api.Walker(
+ pd.DataFrame([{"city": "London"}]),
+ gid="streamlit-core",
+ computation="browser",
+ )
+ renderer = streamlit.StreamlitRenderer(public_walker)
+ html = renderer._get_html(mode="table")
+
+ assert len(FakeWalker.instances) == 1
+ assert renderer.walker is public_walker.core
+ assert public_walker.core.use_preview is False
+ assert public_walker.core.is_export_dataframe is False
+ assert public_walker.core.init_callback_calls == [({"gid": "streamlit-core"}, None)]
+ rendered_props = json.loads(html)
+ assert rendered_props["env"] == "streamlit"
+ assert rendered_props["gwMode"] == "table"
+ component = renderer.explorer(key="public-explorer")
+ assert public_walker.core.props_calls[-1] == ("streamlit", [{"city": "London"}], False)
+ assert component["dataSource"] == [{"city": "London"}]
+
+
+def test_streamlit_renderer_component_data_source_matches_computation_mode(monkeypatch):
+ _install_streamlit_stubs(monkeypatch)
+ streamlit = importlib.reload(importlib.import_module("pygwalker.api.streamlit"))
+
+ _reset_fake_walker()
+ monkeypatch.setattr(streamlit, "PygWalker", FakeWalker)
+ monkeypatch.setattr(streamlit, "init_streamlit_comm", lambda: None)
+ monkeypatch.setattr(streamlit, "get_dataset_hash", lambda _dataset: "dataset-hash")
+ monkeypatch.setattr(streamlit, "StreamlitCommunication", lambda gid: {"gid": gid})
+
+ browser_renderer = streamlit.StreamlitRenderer(
+ pd.DataFrame([{"city": "London"}]),
+ computation="browser",
+ )
+ FakeWalker.instances[-1].origin_data_source = [{"date": pd.Timestamp("2024-01-01"), "city": "London"}]
+ browser_component = browser_renderer.explorer(key="browser-explorer")
+
+ assert FakeWalker.instances[-1].props_calls[-1] == ("streamlit", [{"date": 1704067200000, "city": "London"}], False)
+ assert browser_component["dataSource"] == [{"date": 1704067200000, "city": "London"}]
+
+ kernel_renderer = streamlit.StreamlitRenderer(
+ pd.DataFrame([{"city": "London"}]),
+ computation="kernel",
+ )
+ kernel_component = kernel_renderer.explorer(key="kernel-explorer")
+
+ assert FakeWalker.instances[-1].props_calls[-1] == ("streamlit", [], False)
+ assert kernel_component["dataSource"] == []
+
+
+def test_streamlit_renderer_rejects_rebuilding_public_walker_object(monkeypatch, tmp_path):
+ _install_streamlit_stubs(monkeypatch)
+ streamlit = importlib.reload(importlib.import_module("pygwalker.api.streamlit"))
+ walker_api = importlib.import_module("pygwalker.api.walker")
+
+ _reset_fake_walker()
+ monkeypatch.setattr(walker_api, "PygWalker", FakeWalker)
+ monkeypatch.setattr(streamlit, "init_streamlit_comm", lambda: None)
+
+ public_walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), computation="browser")
+
+ with pytest.raises(ValueError, match="cannot apply construction parameters: spec_path"):
+ streamlit.StreamlitRenderer(public_walker, spec_path=str(tmp_path / "other.json"))
+
+
+def test_streamlit_html_accepts_public_walker_object(monkeypatch):
+ _install_streamlit_stubs(monkeypatch)
+ streamlit = importlib.reload(importlib.import_module("pygwalker.api.streamlit"))
+ walker_api = importlib.import_module("pygwalker.api.walker")
+
+ _reset_fake_walker()
+ monkeypatch.setattr(walker_api, "PygWalker", FakeWalker)
+ monkeypatch.setattr(streamlit, "init_streamlit_comm", lambda: None)
+ monkeypatch.setattr(streamlit, "StreamlitCommunication", lambda gid: {"gid": gid})
+
+ public_walker = walker_api.Walker(
+ pd.DataFrame([{"city": "London"}]),
+ gid="streamlit-html",
+ computation="browser",
+ )
+ html = streamlit.get_streamlit_html(public_walker, mode="table")
+
+ assert len(FakeWalker.instances) == 1
+ rendered_props = json.loads(html)
+ assert rendered_props["env"] == "streamlit"
+ assert rendered_props["gwMode"] == "table"
+
+
+def test_webserver_walk_builds_walker_and_starts_server(monkeypatch):
+ from pygwalker.api import webserver
+
+ _reset_fake_walker()
+ monkeypatch.setattr(webserver, "PygWalker", FakeWalker)
+ starts = []
+ monkeypatch.setattr(
+ webserver,
+ "_start_server",
+ lambda walker, port, *, auto_open, auto_shutdown: starts.append((walker, port, auto_open, auto_shutdown)),
+ )
+
+ webserver.walk(
+ pd.DataFrame([{"city": "London"}]),
+ gid="server",
+ port=8765,
+ auto_open=True,
+ auto_shutdown=True,
+ computation="kernel",
+ )
+
+ walker = FakeWalker.instances[0]
+ assert walker.kwargs["gid"] == "server"
+ assert walker.kwargs["use_save_tool"] is True
+ assert walker.kwargs["is_export_dataframe"] is True
+ assert walker.kwargs["kernel_computation"] is True
+ assert walker.kwargs["cloud_computation"] is False
+ assert walker.kwargs["gw_mode"] == "explore"
+ assert starts == [(walker, 8765, True, True)]
+
+
+def test_webserver_walk_resolves_legacy_use_kernel_calc(monkeypatch):
+ from pygwalker.api import webserver
+
+ _reset_fake_walker()
+ monkeypatch.setattr(webserver, "PygWalker", FakeWalker)
+ starts = []
+ monkeypatch.setattr(
+ webserver,
+ "_start_server",
+ lambda walker, port, *, auto_open, auto_shutdown: starts.append((walker, port, auto_open, auto_shutdown)),
+ )
+
+ with pytest.warns(DeprecationWarning, match="use_kernel_calc"):
+ webserver.walk(
+ pd.DataFrame([{"city": "London"}]),
+ use_kernel_calc=True,
+ )
+
+ walker = FakeWalker.instances[0]
+ assert walker.kwargs["kernel_computation"] is True
+ assert walker.kwargs["cloud_computation"] is False
+ assert starts == [(walker, None, False, False)]
+
+
+def test_webserver_walk_accepts_public_walker_object(monkeypatch):
+ from pygwalker.api import webserver
+ from pygwalker.api import walker as walker_api
+
+ _reset_fake_walker()
+ monkeypatch.setattr(walker_api, "PygWalker", FakeWalker)
+ starts = []
+ monkeypatch.setattr(
+ webserver,
+ "_start_server",
+ lambda walker, port, *, auto_open, auto_shutdown: starts.append((walker, port, auto_open, auto_shutdown)),
+ )
+
+ public_walker = walker_api.Walker(
+ pd.DataFrame([{"city": "London"}]),
+ gid="server-core",
+ computation="browser",
+ )
+ webserver.walk(public_walker, port=8768, auto_open=True, auto_shutdown=True)
+
+ assert len(FakeWalker.instances) == 1
+ assert starts == [(public_walker.core, 8768, True, True)]
+
+
+def test_webserver_walk_rejects_rebuilding_public_walker_object(monkeypatch, tmp_path):
+ from pygwalker.api import webserver
+ from pygwalker.api import walker as walker_api
+
+ _reset_fake_walker()
+ monkeypatch.setattr(walker_api, "PygWalker", FakeWalker)
+ public_walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), computation="browser")
+
+ with pytest.raises(ValueError, match="cannot apply construction parameters: spec_path"):
+ webserver.walk(public_walker, spec_path=str(tmp_path / "other.json"))
+
+
+def test_webserver_walk_allows_none_cloud_computation_for_public_walker(monkeypatch):
+ from pygwalker.api import webserver
+ from pygwalker.api import walker as walker_api
+
+ _reset_fake_walker()
+ monkeypatch.setattr(walker_api, "PygWalker", FakeWalker)
+ starts = []
+ monkeypatch.setattr(
+ webserver,
+ "_start_server",
+ lambda walker, port, *, auto_open, auto_shutdown: starts.append((walker, port, auto_open, auto_shutdown)),
+ )
+
+ public_walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), computation="browser")
+ webserver.walk(public_walker, cloud_computation=None)
+
+ assert starts == [(public_walker.core, None, False, False)]
+
+
+def test_webserver_walk_rejects_show_cloud_tool_true_alias_for_public_walker(monkeypatch):
+ from pygwalker.api import webserver
+ from pygwalker.api import walker as walker_api
+
+ _reset_fake_walker()
+ monkeypatch.setattr(walker_api, "PygWalker", FakeWalker)
+ public_walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), computation="browser")
+
+ with pytest.raises(ValueError, match="cannot apply construction parameters: show_cloud_tool"):
+ webserver.walk(public_walker, show_cloud_tool=1)
+
+
+def test_webserver_walk_rejects_legacy_kernel_flag_for_public_walker(monkeypatch):
+ from pygwalker.api import webserver
+ from pygwalker.api import walker as walker_api
+
+ _reset_fake_walker()
+ monkeypatch.setattr(walker_api, "PygWalker", FakeWalker)
+ public_walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), computation="browser")
+
+ with pytest.raises(ValueError, match="cannot apply construction parameters: use_kernel_calc"):
+ webserver.walk(public_walker, use_kernel_calc=True)
+
+
+def test_webserver_start_server_disables_preview(monkeypatch):
+ from pygwalker.api import webserver
+
+ class FakeServer:
+ def __init__(self, *_args, **_kwargs):
+ pass
+
+ def __enter__(self):
+ return self
+
+ def __exit__(self, *_args):
+ return False
+
+ def serve_forever(self):
+ return None
+
+ walker = SimpleNamespace(
+ gid="server-preview",
+ use_preview=True,
+ init_callback_calls=[],
+ )
+ walker._init_callback = lambda comm, preview_tool=None: walker.init_callback_calls.append((comm, preview_tool))
+
+ monkeypatch.setattr(webserver, "CustomTCPServer", FakeServer)
+
+ webserver._start_server(walker, 8769, auto_open=False, auto_shutdown=False)
+
+ assert walker.use_preview is False
+ assert len(walker.init_callback_calls) == 1
+ assert walker.init_callback_calls[0][1] is None
+
+
+def test_webserver_render_builds_filter_renderer(monkeypatch):
+ from pygwalker.api import webserver
+
+ _reset_fake_walker()
+ monkeypatch.setattr(webserver, "PygWalker", FakeWalker)
+ starts = []
+ monkeypatch.setattr(
+ webserver,
+ "_start_server",
+ lambda walker, port, *, auto_open, auto_shutdown: starts.append((walker, port, auto_open, auto_shutdown)),
+ )
+
+ webserver.render(
+ pd.DataFrame([{"city": "London"}]),
+ spec="{}",
+ port=8766,
+ auto_open=False,
+ auto_shutdown=True,
+ computation="browser",
+ )
+
+ walker = FakeWalker.instances[0]
+ assert walker.kwargs["spec"] == "{}"
+ assert walker.kwargs["use_save_tool"] is False
+ assert walker.kwargs["kernel_computation"] is False
+ assert walker.kwargs["cloud_computation"] is False
+ assert walker.kwargs["gw_mode"] == "filter_renderer"
+ assert walker.kwargs["is_export_dataframe"] is True
+ assert starts == [(walker, 8766, False, True)]
+
+
+def test_webserver_table_builds_table_renderer(monkeypatch):
+ from pygwalker.api import webserver
+
+ _reset_fake_walker()
+ monkeypatch.setattr(webserver, "PygWalker", FakeWalker)
+ starts = []
+ monkeypatch.setattr(
+ webserver,
+ "_start_server",
+ lambda walker, port, *, auto_open, auto_shutdown: starts.append((walker, port, auto_open, auto_shutdown)),
+ )
+
+ with pytest.warns(DeprecationWarning, match="kernel_computation"):
+ webserver.table(
+ pd.DataFrame([{"city": "London"}]),
+ port=8767,
+ auto_open=True,
+ auto_shutdown=False,
+ kernel_computation=False,
+ )
+
+ walker = FakeWalker.instances[0]
+ assert walker.kwargs["spec"] == ""
+ assert walker.kwargs["use_save_tool"] is False
+ assert walker.kwargs["gw_mode"] == "table"
+ assert walker.kwargs["is_export_dataframe"] is True
+ assert starts == [(walker, 8767, True, False)]
+
+
+def _install_anywidget_stubs(monkeypatch):
+ import traitlets
+
+ class FakeAnyWidget(traitlets.HasTraits):
+ pass
+
+ monkeypatch.setitem(sys.modules, "anywidget", SimpleNamespace(AnyWidget=FakeAnyWidget))
+
+
+def _install_streamlit_stubs(monkeypatch):
+ fake_streamlit = ModuleType("streamlit")
+ fake_streamlit.config = SimpleNamespace(get_option=lambda _name: "")
+ fake_streamlit.cache_resource = lambda func: func
+
+ fake_components = ModuleType("streamlit.components")
+ fake_components_v1 = ModuleType("streamlit.components.v1")
+ fake_components_v1.declare_component = lambda *_args, **_kwargs: lambda **component_kwargs: component_kwargs
+ fake_components_module = ModuleType("streamlit.components.v1.components")
+ fake_components_module.CustomComponent = object
+
+ fake_streamlit.components = fake_components
+ fake_components.v1 = fake_components_v1
+
+ monkeypatch.setitem(sys.modules, "streamlit", fake_streamlit)
+ monkeypatch.setitem(sys.modules, "streamlit.components", fake_components)
+ monkeypatch.setitem(sys.modules, "streamlit.components.v1", fake_components_v1)
+ monkeypatch.setitem(sys.modules, "streamlit.components.v1.components", fake_components_module)
diff --git a/tests/test_metrics_tools.py b/tests/test_metrics_tools.py
new file mode 100644
index 00000000..d70a1573
--- /dev/null
+++ b/tests/test_metrics_tools.py
@@ -0,0 +1,62 @@
+import pandas as pd
+import pytest
+
+from pygwalker_tools.metrics import get_metrics_datas
+from pygwalker_tools.metrics.core import get_metrics_sql
+
+
+def _event_dataframe():
+ return pd.DataFrame(
+ [
+ {"date": "2024-01-01", "user_id": "alice", "user_signup_date": "2024-01-01"},
+ {"date": "2024-01-01", "user_id": "bob", "user_signup_date": "2024-01-01"},
+ {"date": "2024-01-02", "user_id": "alice", "user_signup_date": "2024-01-01"},
+ ]
+ )
+
+
+def test_metrics_sql_maps_fields_into_subquery():
+ sql = get_metrics_sql(
+ name="pv",
+ field_map={"date": "event_date"},
+ params={},
+ origin_table_name="pygwalker_mid_table",
+ )
+
+ assert 'CAST("event_date" AS TIMESTAMP) AS "date"' in sql
+ assert 'FROM "pygwalker_mid_table"' in sql
+ assert 'COUNT(1) AS "pv"' in sql
+
+
+def test_metrics_sql_validates_metric_fields_and_params():
+ with pytest.raises(ValueError, match="Unknown metrics name: missing"):
+ get_metrics_sql(name="missing", field_map={}, params={}, origin_table_name="source")
+
+ with pytest.raises(ValueError, match="Field not found: user_id"):
+ get_metrics_sql(name="uv", field_map={"date": "date"}, params={}, origin_table_name="source")
+
+ with pytest.raises(ValueError, match="Param not found: time_unit"):
+ get_metrics_sql(
+ name="retention",
+ field_map={"date": "date", "user_id": "user_id", "user_signup_date": "signup_date"},
+ params={},
+ origin_table_name="source",
+ )
+
+
+def test_get_metrics_datas_queries_pandas_dataset():
+ datas = get_metrics_datas(
+ _event_dataframe(),
+ "uv",
+ {"date": "date", "user_id": "user_id"},
+ )
+
+ assert sorted(datas, key=lambda row: row["date"]) == [
+ {"date": "2024-01-01", "uv": 2},
+ {"date": "2024-01-02", "uv": 1},
+ ]
+
+
+def test_get_metrics_datas_rejects_cloud_dataset_id():
+ with pytest.raises(TypeError, match="Unsupported cloud dataset type"):
+ get_metrics_datas("cloud-dataset-id", "pv", {"date": "date"})
diff --git a/tests/test_preview_image.py b/tests/test_preview_image.py
new file mode 100644
index 00000000..292bf9ce
--- /dev/null
+++ b/tests/test_preview_image.py
@@ -0,0 +1,61 @@
+from pygwalker.services import preview_image
+
+
+def test_preview_image_tool_renders_on_daemon_thread(monkeypatch):
+ started_threads = []
+ display_calls = []
+
+ class FakeThread:
+ def __init__(self, *, target, daemon):
+ self.target = target
+ self.daemon = daemon
+
+ def start(self):
+ started_threads.append(self)
+
+ monkeypatch.setattr(preview_image, "Thread", FakeThread)
+ monkeypatch.setattr(preview_image, "_create_jupyter_frontend", lambda: None)
+ monkeypatch.setattr(preview_image, "display_html", lambda html, slot_id: display_calls.append((html, slot_id)))
+
+ tool = preview_image.PreviewImageTool("preview")
+ tool.async_render_gw_review("chart
")
+ tool.async_render_gw_review("chart 2
")
+
+ assert len(started_threads) == 1
+ assert started_threads[0].daemon is True
+ tool._render_next_preview()
+ tool._render_next_preview()
+ assert display_calls == [
+ ("chart
", "pygwalker-preview-preview"),
+ ("chart 2
", "pygwalker-preview-preview"),
+ ]
+
+
+def test_preview_image_tool_keeps_worker_alive_after_render_error(monkeypatch):
+ display_calls = []
+
+ class FakeThread:
+ def __init__(self, *, target, daemon):
+ self.target = target
+ self.daemon = daemon
+
+ def start(self):
+ pass
+
+ def fake_display_html(html, slot_id):
+ if html == "bad chart
":
+ raise RuntimeError("render failed")
+ display_calls.append((html, slot_id))
+
+ monkeypatch.setattr(preview_image, "Thread", FakeThread)
+ monkeypatch.setattr(preview_image, "_create_jupyter_frontend", lambda: None)
+ monkeypatch.setattr(preview_image, "display_html", fake_display_html)
+
+ tool = preview_image.PreviewImageTool("preview")
+ tool.async_render_gw_review("bad chart
")
+ tool.async_render_gw_review("good chart
")
+
+ tool._safe_render_next_preview()
+ tool._safe_render_next_preview()
+
+ assert display_calls == [("good chart
", "pygwalker-preview-preview")]
diff --git a/tests/test_privacy_notice.py b/tests/test_privacy_notice.py
new file mode 100644
index 00000000..bf52a108
--- /dev/null
+++ b/tests/test_privacy_notice.py
@@ -0,0 +1,169 @@
+import subprocess
+import sys
+from types import SimpleNamespace
+
+from pygwalker.services import config, track
+from pygwalker.services.global_var import GlobalVarManager
+
+
+def test_default_privacy_is_update_only():
+ assert config.DEFAULT_CONFIG["privacy"] == "update-only"
+ assert config.privacy_item.default == "update-only"
+
+
+def test_privacy_notice_sentinel_is_written_once(monkeypatch, tmp_path):
+ notice_path = tmp_path / "privacy_notice_shown"
+ monkeypatch.setattr(config, "PRIVACY_NOTICE_PATH", str(notice_path))
+
+ assert config.should_show_privacy_notice() is True
+ assert notice_path.read_text() == "shown"
+ assert config.should_show_privacy_notice() is False
+
+
+def test_track_event_prints_privacy_notice_once_and_handles_empty_properties(monkeypatch, capsys):
+ analytics_calls = []
+ kanaries_calls = []
+ previous_privacy = GlobalVarManager.privacy
+ GlobalVarManager.privacy = "events"
+
+ monkeypatch.setattr(track, "should_show_privacy_notice", lambda: True)
+ monkeypatch.setattr(track, "get_local_user_id", lambda: "test-user")
+ monkeypatch.setattr(
+ track,
+ "_get_analytics_client",
+ lambda: SimpleNamespace(track=lambda **kwargs: analytics_calls.append(kwargs)),
+ )
+ monkeypatch.setattr(
+ track,
+ "_get_kanaries_track_client",
+ lambda: SimpleNamespace(track=lambda payload: kanaries_calls.append(payload)),
+ )
+
+ try:
+ track.track_event("invoke_props")
+ finally:
+ GlobalVarManager.privacy = previous_privacy
+
+ assert "pygwalker config --set privacy=update-only" in capsys.readouterr().out
+ assert analytics_calls == [{"user_id": "test-user", "event": "invoke_props", "properties": {}}]
+ assert kanaries_calls == [{"user_id": "test-user"}]
+
+
+def test_track_event_does_not_emit_when_privacy_is_not_events(monkeypatch, capsys):
+ notice_calls = []
+ analytics_calls = []
+ previous_privacy = GlobalVarManager.privacy
+ GlobalVarManager.privacy = "update-only"
+
+ monkeypatch.setattr(track, "should_show_privacy_notice", lambda: notice_calls.append(True) or True)
+ monkeypatch.setattr(
+ track,
+ "_get_analytics_client",
+ lambda: SimpleNamespace(track=lambda **kwargs: analytics_calls.append(kwargs)),
+ )
+
+ try:
+ track.track_event("invoke_props", {"mode": "test"})
+ finally:
+ GlobalVarManager.privacy = previous_privacy
+
+ assert capsys.readouterr().out == ""
+ assert notice_calls == []
+ assert analytics_calls == []
+
+
+def test_track_event_does_not_import_analytics_clients_when_privacy_is_not_events():
+ code = """
+import builtins
+
+original_import = builtins.__import__
+
+def blocked_import(name, *args, **kwargs):
+ if name == "kanaries_track" or name == "segment" or name.startswith("segment."):
+ raise AssertionError(f"unexpected analytics import: {name}")
+ return original_import(name, *args, **kwargs)
+
+builtins.__import__ = blocked_import
+
+from pygwalker.services.global_var import GlobalVarManager
+from pygwalker.services.track import track_event
+
+GlobalVarManager.privacy = "update-only"
+track_event("invoke_props", {"mode": "test"})
+"""
+ result = subprocess.run(
+ [sys.executable, "-c", code],
+ check=False,
+ text=True,
+ capture_output=True,
+ )
+
+ assert result.returncode == 0, result.stderr
+
+
+def test_analytics_clients_are_configured_without_background_workers(monkeypatch):
+ fake_analytics = SimpleNamespace(
+ write_key=None,
+ sync_mode=False,
+ timeout=15,
+ max_retries=10,
+ track=lambda **_kwargs: None,
+ )
+ fake_kanaries_track = SimpleNamespace(
+ config=SimpleNamespace(
+ auth_token=None,
+ proxies=None,
+ sync_send=False,
+ timeout=15,
+ max_retries=5,
+ thread=1,
+ ),
+ track=lambda _payload: None,
+ )
+
+ previous_analytics_client = track._analytics_client
+ previous_kanaries_track_client = track._kanaries_track_client
+ monkeypatch.setitem(sys.modules, "segment", SimpleNamespace(analytics=fake_analytics))
+ monkeypatch.setitem(sys.modules, "segment.analytics", fake_analytics)
+ monkeypatch.setitem(sys.modules, "kanaries_track", fake_kanaries_track)
+ track._analytics_client = None
+ track._kanaries_track_client = None
+
+ try:
+ assert track._get_analytics_client() is fake_analytics
+ assert track._get_kanaries_track_client() is fake_kanaries_track
+ finally:
+ track._analytics_client = previous_analytics_client
+ track._kanaries_track_client = previous_kanaries_track_client
+
+ assert fake_analytics.sync_mode is True
+ assert fake_analytics.timeout == 1
+ assert fake_analytics.max_retries == 0
+ assert fake_kanaries_track.config.sync_send is True
+ assert fake_kanaries_track.config.timeout == 1
+ assert fake_kanaries_track.config.max_retries == 1
+ assert fake_kanaries_track.config.thread == 0
+
+
+def test_kanaries_track_single_retry_setting_does_not_loop(monkeypatch):
+ from kanaries_track.request import RequestClient
+
+ calls = []
+ client = RequestClient(
+ host="https://example.invalid",
+ auth_token="test-token",
+ max_retries=1,
+ timeout=1,
+ verify=True,
+ proxy={},
+ )
+
+ def fail_post(*_args, **_kwargs):
+ calls.append(True)
+ raise OSError("network unavailable")
+
+ monkeypatch.setattr(client.session, "post", fail_post)
+
+ client.track([{"event": "test"}])
+
+ assert calls == [True]
diff --git a/tests/test_public_spec_api.py b/tests/test_public_spec_api.py
new file mode 100644
index 00000000..5b5b62a6
--- /dev/null
+++ b/tests/test_public_spec_api.py
@@ -0,0 +1,105 @@
+import json
+
+import pytest
+
+import pygwalker as pyg
+from pygwalker import __version__
+from pygwalker.services import spec as spec_service
+
+
+def _legacy_spec():
+ return {
+ "config": [
+ {
+ "name": "Legacy chart",
+ "config": {},
+ "encodings": {
+ "dimensions": [
+ {
+ "fid": "date",
+ "name": "date",
+ "computed": True,
+ "expression": {
+ "params": [
+ {"type": "field", "value": "date"},
+ {"type": "offset", "value": -480},
+ ]
+ },
+ }
+ ],
+ "measures": [],
+ },
+ "visId": "legacy",
+ }
+ ],
+ "chart_map": {},
+ "workflow_list": [{"workflow": []}],
+ "version": "0.4.7a5",
+ }
+
+
+def test_spec_migrate_updates_legacy_spec_json_to_current_version():
+ migrated = pyg.spec.migrate(json.dumps(_legacy_spec()))
+
+ chart = migrated["config"][0]
+ field = chart["encodings"]["dimensions"][0]
+
+ assert migrated["version"] == __version__
+ assert migrated["workflow_list"] == [{"workflow": []}]
+ assert chart["config"]["timezoneDisplayOffset"] == 0
+ assert field["offset"] == 0
+ assert field["expression"]["params"][1] == {"type": "offset", "value": 0}
+
+
+def test_spec_migrate_accepts_existing_local_file_and_custom_version(tmp_path):
+ spec_path = tmp_path / "legacy.json"
+ spec_path.write_text(json.dumps(_legacy_spec()), encoding="utf-8")
+
+ migrated = pyg.spec.migrate(str(spec_path), version="0.6.0")
+
+ assert migrated["version"] == "0.6.0"
+ assert migrated["config"][0]["name"] == "Legacy chart"
+
+
+def test_spec_migrate_reads_32_char_hex_local_file_without_remote_lookup(tmp_path, monkeypatch):
+ spec_path = tmp_path / "0123456789abcdef0123456789abcdef"
+ spec_path.write_text(json.dumps(_legacy_spec()), encoding="utf-8")
+
+ def fail_remote_lookup(config_id):
+ raise AssertionError(f"unexpected remote lookup for {config_id}")
+
+ monkeypatch.setattr(spec_service, "_get_spec_from_server", fail_remote_lookup)
+
+ migrated = pyg.spec.migrate(str(spec_path))
+
+ assert migrated["config"][0]["name"] == "Legacy chart"
+
+
+def test_spec_migrate_rejects_missing_path_without_creating_file(tmp_path):
+ missing_path = tmp_path / "missing.json"
+
+ with pytest.raises(ValueError, match="dict, list, JSON string, or existing local file path"):
+ pyg.spec.migrate(str(missing_path))
+
+ assert not missing_path.exists()
+
+
+def test_spec_migrate_rejects_invalid_json_string():
+ with pytest.raises(ValueError, match="spec is not a valid json"):
+ pyg.spec.migrate("{not-valid-json")
+
+
+def test_spec_migrate_rejects_json_scalar():
+ with pytest.raises(ValueError, match="JSON object or array"):
+ pyg.spec.migrate("42")
+
+
+def test_spec_migrate_does_not_mutate_caller_spec():
+ legacy = _legacy_spec()
+ original = json.loads(json.dumps(legacy))
+
+ migrated = pyg.spec.migrate(legacy)
+
+ assert legacy == original
+ assert migrated is not legacy
+ assert migrated["config"][0]["config"]["timezoneDisplayOffset"] == 0
diff --git a/tests/test_pydantic_compat.py b/tests/test_pydantic_compat.py
new file mode 100644
index 00000000..af06df35
--- /dev/null
+++ b/tests/test_pydantic_compat.py
@@ -0,0 +1,20 @@
+from pathlib import Path
+
+
+PYDANTIC_V1_API_PATTERNS = (".dict(", ".parse_obj(")
+PYDANTIC_COMPAT_MODULE = Path("pygwalker/utils/pydantic_compat.py")
+
+
+def test_pydantic_v1_api_calls_stay_in_compat_module():
+ repo_root = Path(__file__).resolve().parents[1]
+ offenders = []
+
+ for path in sorted((repo_root / "pygwalker").rglob("*.py")):
+ relative_path = path.relative_to(repo_root)
+ if relative_path == PYDANTIC_COMPAT_MODULE:
+ continue
+ source = path.read_text(encoding="utf-8")
+ if any(pattern in source for pattern in PYDANTIC_V1_API_PATTERNS):
+ offenders.append(str(relative_path))
+
+ assert offenders == []
diff --git a/tests/test_pygwalker_core.py b/tests/test_pygwalker_core.py
index 28ec42b4..144a2434 100644
--- a/tests/test_pygwalker_core.py
+++ b/tests/test_pygwalker_core.py
@@ -1,15 +1,41 @@
+import base64
+from contextlib import nullcontext
+import json
from types import SimpleNamespace
+import urllib.parse
+import zlib
import pandas as pd
+import pyarrow as pa
import pytest
+from duckdb import ParserException
+from pygwalker import __version__
from pygwalker.api import adapter, html, jupyter
+from pygwalker.api.component import Component
from pygwalker.api import pygwalker as pygwalker_module
from pygwalker.api.pygwalker import PygWalker
from pygwalker.communications.base import BaseCommunication
+from pygwalker.errors import ErrorCode
+from pygwalker.services import chart_export as chart_export_module
+from pygwalker.services import data_bridge as data_bridge_module
+from pygwalker.services import desktop_import as desktop_import_module
+from pygwalker.services import jupyter_display as jupyter_display_module
+from pygwalker.services import props_tracker as props_tracker_module
+from pygwalker.services import render_manager as render_manager_module
from pygwalker.services.global_var import GlobalVarManager
+def _expected_legacy_computation_warning(kwargs):
+ if (
+ kwargs.get("kernel_computation") is not None
+ or kwargs.get("use_kernel_calc") is not None
+ or kwargs.get("cloud_computation") is True
+ ):
+ return pytest.warns(DeprecationWarning, match="deprecated")
+ return nullcontext()
+
+
def _make_walker(monkeypatch, **kwargs):
monkeypatch.setattr(pygwalker_module, "check_update", lambda: None)
monkeypatch.setattr(pygwalker_module, "track_event", lambda *_args, **_kwargs: None)
@@ -42,6 +68,74 @@ def _make_walker(monkeypatch, **kwargs):
return PygWalker(**defaults)
+def _patch_cloud_computation_parser(monkeypatch):
+ original_get_parser = data_bridge_module.get_parser
+ uploaded = []
+
+ class FakeCloudParser:
+ data_size = 1
+ raw_fields = [
+ {"fid": "city", "name": "city", "semanticType": "nominal", "analyticType": "dimension"},
+ {"fid": "value", "name": "value", "semanticType": "quantitative", "analyticType": "measure"},
+ ]
+ dataset_type = "cloud_dataset"
+
+ def to_records(self, limit=None):
+ return [{"city": "London", "value": 1}]
+
+ def fake_get_parser(
+ dataset,
+ field_specs=None,
+ infer_string_to_date=False,
+ infer_number_to_dimension=True,
+ other_params=None,
+ ):
+ if isinstance(dataset, str) and dataset == "cloud-dataset-id":
+ return FakeCloudParser()
+ return original_get_parser(
+ dataset,
+ field_specs,
+ infer_string_to_date,
+ infer_number_to_dimension,
+ other_params,
+ )
+
+ def fake_create_cloud_dataset(self, data_parser, name, is_public, temporary):
+ uploaded.append(
+ {
+ "data_parser": data_parser,
+ "name": name,
+ "is_public": is_public,
+ "temporary": temporary,
+ }
+ )
+ return "cloud-dataset-id"
+
+ monkeypatch.setattr(data_bridge_module, "get_parser", fake_get_parser)
+ monkeypatch.setattr(pygwalker_module.CloudService, "create_cloud_dataset", fake_create_cloud_dataset)
+ return uploaded
+
+
+def _chart_payload(title="Updated chart"):
+ return {
+ "charts": [
+ {
+ "rowIndex": 0,
+ "colIndex": 0,
+ "data": "data:image/png;base64,abc",
+ "height": 100,
+ "width": 200,
+ "canvasHeight": 100,
+ "canvasWidth": 200,
+ }
+ ],
+ "singleChart": "data:image/png;base64,abc",
+ "nRows": 1,
+ "nCols": 1,
+ "title": title,
+ }
+
+
def test_pygwalker_props_expose_browser_data_path(monkeypatch):
walker = _make_walker(monkeypatch, kernel_computation=False)
@@ -72,6 +166,357 @@ def test_pygwalker_props_expose_kernel_data_path(monkeypatch):
assert props["useKernelCalc"] is True
+def test_pygwalker_get_props_falls_back_without_props_builder(monkeypatch):
+ track_calls = []
+ monkeypatch.setattr(pygwalker_module, "get_local_user_id", lambda: "fallback-user")
+ monkeypatch.setattr(pygwalker_module, "track_event", lambda *args, **kwargs: track_calls.append((args, kwargs)))
+ walker = PygWalker.__new__(PygWalker)
+ walker.gid = "fallback"
+ walker.data_bridge = SimpleNamespace(
+ origin_data_source=[{"city": "London"}],
+ field_specs=[{"fid": "city"}],
+ kernel_computation=False,
+ parse_dsl_type="client",
+ dataset_type="pandas_dataframe",
+ data_parser=SimpleNamespace(field_metas=[]),
+ )
+ walker.spec_manager = SimpleNamespace(
+ vis_spec=[],
+ spec_type="empty",
+ chart_map={},
+ )
+ walker.theme_key = "g2"
+ walker.appearance = "light"
+ walker.source_invoke_code = ""
+ walker.tunnel_id = "tunnel"
+ walker.data_source_id = "data-source"
+ walker.show_cloud_tool = False
+ walker.use_save_tool = False
+ walker.gw_mode = "explore"
+ walker.other_props = {}
+ walker.is_export_dataframe = False
+ walker.default_tab = "vis"
+ walker.cloud_computation = False
+ walker.kanaries_api_key = ""
+
+ props = walker._get_props("fallback-env")
+
+ assert props["id"] == "fallback"
+ assert props["hashcode"] == "fallback-user"
+ assert props["dataSource"] == [{"city": "London"}]
+ assert props["rawFields"] == [{"fid": "city", "offset": 0}]
+ assert track_calls[0][0][0] == "invoke_props"
+
+
+def test_props_tracker_tracks_expected_invocation_fields():
+ track_calls = []
+ tracker = props_tracker_module.PropsTracker(
+ SimpleNamespace(kanaries_api_key="secret-token"),
+ lambda event, props: track_calls.append((event, props)),
+ )
+
+ tracker.track_invocation(
+ {
+ "id": "core",
+ "version": "0.5",
+ "hashcode": "user",
+ "themeKey": "g2",
+ "dark": "light",
+ "env": "jupyter",
+ "specType": "empty",
+ "needLoadDatas": True,
+ "showCloudTool": False,
+ "useKernelCalc": False,
+ "useSaveTool": True,
+ "parseDslType": "client",
+ "gwMode": "explore",
+ "datasetType": "pandas_dataframe",
+ "defaultTab": "vis",
+ "useCloudCalc": False,
+ "dataSource": [{"city": "London"}],
+ "rawFields": [{"fid": "city"}],
+ }
+ )
+
+ assert track_calls == [
+ (
+ "invoke_props",
+ {
+ "id": "core",
+ "version": "0.5",
+ "hashcode": "user",
+ "themeKey": "g2",
+ "dark": "light",
+ "env": "jupyter",
+ "specType": "empty",
+ "needLoadDatas": True,
+ "showCloudTool": False,
+ "useKernelCalc": False,
+ "useSaveTool": True,
+ "parseDslType": "client",
+ "gwMode": "explore",
+ "datasetType": "pandas_dataframe",
+ "defaultTab": "vis",
+ "useCloudCalc": False,
+ "hasKanariesToken": True,
+ },
+ )
+ ]
+
+
+def test_render_manager_preview_handles_parser_exception_and_manual_gid(monkeypatch):
+ captured = {}
+
+ class FakeDataParser:
+ def get_datas_by_payload(self, workflow):
+ if workflow == "bad-workflow":
+ raise ParserException("bad workflow")
+ return [{"workflow": workflow}]
+
+ def fake_render_preview(vis_spec, datas, theme_key, gid, appearance):
+ captured.update(
+ {
+ "vis_spec": vis_spec,
+ "datas": datas,
+ "theme_key": theme_key,
+ "gid": gid,
+ "appearance": appearance,
+ }
+ )
+ return "preview-html"
+
+ monkeypatch.setattr(render_manager_module, "render_gw_preview_html", fake_render_preview)
+ monkeypatch.setattr(render_manager_module, "rand_str", lambda: "-manual")
+
+ walker = SimpleNamespace(
+ workflow_list=["good-workflow", "bad-workflow"],
+ data_parser=FakeDataParser(),
+ vis_spec=[{"chart": "bar"}],
+ theme_key="g2",
+ gid="core",
+ appearance="light",
+ )
+
+ html = render_manager_module.RenderManager(walker).get_preview_html(manual=True)
+
+ assert html == "preview-html"
+ assert captured == {
+ "vis_spec": [{"chart": "bar"}],
+ "datas": [[{"workflow": "good-workflow"}], []],
+ "theme_key": "g2",
+ "gid": "core-manual",
+ "appearance": "light",
+ }
+
+
+def test_render_manager_chart_preview_uses_chart_indexed_workflow(monkeypatch):
+ captured = {}
+
+ class FakeDataParser:
+ def get_datas_by_payload(self, workflow):
+ assert workflow == "target-workflow"
+ return [{"value": 1}]
+
+ class FakeSpecManager:
+ def get_chart_index(self, chart_name):
+ assert chart_name == "Target chart"
+ return 1
+
+ def fake_render_chart_preview(**kwargs):
+ captured.update(kwargs)
+ return "chart-html"
+
+ monkeypatch.setattr(render_manager_module, "render_gw_chart_preview_html", fake_render_chart_preview)
+
+ walker = SimpleNamespace(
+ workflow_list=["other-workflow", "target-workflow"],
+ data_parser=FakeDataParser(),
+ spec_manager=FakeSpecManager(),
+ vis_spec=[{"chart": "other"}, {"chart": "target"}],
+ theme_key="g2",
+ appearance="light",
+ )
+
+ html = render_manager_module.RenderManager(walker).get_chart_preview_html(
+ "Target chart",
+ title="Chart title",
+ desc="Chart desc",
+ )
+
+ assert html == "chart-html"
+ assert captured == {
+ "single_vis_spec": {"chart": "target"},
+ "data": [{"value": 1}],
+ "theme_key": "g2",
+ "title": "Chart title",
+ "desc": "Chart desc",
+ "appearance": "light",
+ }
+
+
+def test_render_manager_chart_preview_returns_empty_for_mismatched_workflow_list():
+ class FakeSpecManager:
+ def get_chart_index(self, chart_name):
+ assert chart_name == "Missing workflow"
+ return 1
+
+ walker = SimpleNamespace(
+ workflow_list=["only-workflow"],
+ data_parser=SimpleNamespace(get_datas_by_payload=lambda workflow: [{"value": 1}]),
+ spec_manager=FakeSpecManager(),
+ vis_spec=[{"chart": "only"}, {"chart": "missing-workflow"}],
+ theme_key="g2",
+ appearance="light",
+ )
+
+ assert render_manager_module.RenderManager(walker).get_chart_preview_html("Missing workflow", "Title", "Desc") == ""
+
+
+def test_jupyter_display_manager_convert_html_displays_iframe():
+ displayed = []
+ walker = SimpleNamespace(
+ _get_props=lambda env: {"env": env},
+ _get_render_iframe=lambda props: f"iframe-{props['env']}",
+ )
+
+ jupyter_display_module.JupyterDisplayManager(walker, displayed.append).display_on_convert_html()
+
+ assert displayed == ["iframe-jupyter"]
+
+
+def test_jupyter_display_manager_uploads_large_classic_jupyter_data(monkeypatch):
+ displayed = []
+ upload_calls = []
+
+ class FakeUploadTool:
+ def run(self, **kwargs):
+ upload_calls.append(kwargs)
+
+ monkeypatch.setattr(jupyter_display_module, "get_max_limited_datas", lambda records, _limit: records[:1])
+ monkeypatch.setattr(jupyter_display_module, "BatchUploadDatasToolOnJupyter", lambda: FakeUploadTool())
+ monkeypatch.setattr(jupyter_display_module, "render_iframe_messages_html", lambda gid: f"messages-{gid}")
+
+ walker = SimpleNamespace(
+ gid="classic",
+ origin_data_source=[{"city": "London"}, {"city": "Tokyo"}],
+ data_source_id="data-source",
+ tunnel_id="tunnel",
+ _get_props=lambda env, data_source, need_load_datas: {
+ "env": env,
+ "dataSource": data_source,
+ "needLoadDatas": need_load_datas,
+ },
+ _get_render_iframe=lambda props: f"iframe-{props['env']}-{props['needLoadDatas']}",
+ )
+
+ jupyter_display_module.JupyterDisplayManager(walker, displayed.append).display_on_jupyter()
+
+ assert displayed == ["iframe-jupyter-True", "messages-classic"]
+ assert upload_calls == [
+ {
+ "records": [{"city": "London"}, {"city": "Tokyo"}],
+ "sample_data_count": 0,
+ "data_source_id": "data-source",
+ "gid": "classic",
+ "tunnel_id": "tunnel",
+ }
+ ]
+
+
+def test_chart_export_manager_exports_png(monkeypatch):
+ class FakeResponse:
+ def __enter__(self):
+ return self
+
+ def __exit__(self, *_):
+ return False
+
+ def read(self):
+ return b"png-bytes"
+
+ monkeypatch.setattr(chart_export_module.urllib.request, "urlopen", lambda url: FakeResponse())
+ walker = SimpleNamespace(
+ _get_chart_by_name=lambda name: SimpleNamespace(single_chart=f"https://example.test/{name}.png")
+ )
+
+ assert chart_export_module.ChartExportManager(walker, lambda _: None).export_chart_png("Chart 1") == b"png-bytes"
+
+
+def test_chart_export_manager_exports_svg_variants():
+ encoded_svg = "data:image/svg+xml;base64,PHN2Zz48L3N2Zz4="
+ raw_svg = ""
+
+ encoded_walker = SimpleNamespace(
+ _get_chart_by_name=lambda _name: SimpleNamespace(charts=[SimpleNamespace(data=encoded_svg)])
+ )
+ raw_walker = SimpleNamespace(
+ _get_chart_by_name=lambda _name: SimpleNamespace(charts=[SimpleNamespace(data=raw_svg)])
+ )
+
+ assert (
+ chart_export_module.ChartExportManager(encoded_walker, lambda _: None).export_chart_svg("Chart 1")
+ == b""
+ )
+ assert (
+ chart_export_module.ChartExportManager(raw_walker, lambda _: None).export_chart_svg("Chart 1") == b""
+ )
+
+
+def test_chart_export_manager_display_chart_uses_chart_name_as_default_title():
+ displayed = []
+ calls = []
+ walker = SimpleNamespace(
+ _get_gw_chart_preview_html=lambda chart_name, title, desc: calls.append((chart_name, title, desc)) or "html"
+ )
+
+ chart_export_module.ChartExportManager(walker, displayed.append).display_chart("Chart 1")
+
+ assert calls == [("Chart 1", "Chart 1", "")]
+ assert displayed == ["html"]
+
+
+def test_chart_export_manager_get_single_chart_html_by_spec(monkeypatch):
+ calls = []
+
+ monkeypatch.setattr(chart_export_module, "dsl_to_workflow", lambda spec: calls.append(("workflow", spec)) or "wf")
+ monkeypatch.setattr(
+ chart_export_module,
+ "render_gw_chart_preview_html",
+ lambda **kwargs: calls.append(("render", kwargs)) or "chart-html",
+ )
+
+ walker = SimpleNamespace(
+ data_parser=SimpleNamespace(
+ get_datas_by_payload=lambda workflow: calls.append(("data", workflow)) or [{"x": 1}]
+ ),
+ theme_key="g2",
+ appearance="light",
+ )
+
+ html = chart_export_module.ChartExportManager(walker, lambda _: None).get_single_chart_html_by_spec(
+ spec={"mark": "bar"},
+ title="Title",
+ desc="Desc",
+ )
+
+ assert html == "chart-html"
+ assert calls == [
+ ("workflow", {"mark": "bar"}),
+ ("data", "wf"),
+ (
+ "render",
+ {
+ "single_vis_spec": {"mark": "bar"},
+ "data": [{"x": 1}],
+ "theme_key": "g2",
+ "title": "Title",
+ "desc": "Desc",
+ "appearance": "light",
+ },
+ ),
+ ]
+
+
@pytest.mark.parametrize(
("dataset_type", "expected_parse_dsl_type"),
[
@@ -90,6 +535,67 @@ def test_pygwalker_parse_dsl_type_tracks_dataset_location(
assert walker._get_parse_dsl_type(data_parser) == expected_parse_dsl_type
+def test_pygwalker_data_bridge_property_setters_remain_writable(monkeypatch):
+ walker = _make_walker(monkeypatch)
+ data_parser = SimpleNamespace(dataset_type="custom")
+
+ walker.data_parser = data_parser
+ walker.kernel_computation = True
+ walker.origin_data_source = [{"city": "Paris"}]
+ walker.field_specs = [{"fid": "city"}]
+ walker.parse_dsl_type = "server"
+ walker.dataset_type = "custom_dataset"
+
+ assert walker.data_bridge.data_parser is data_parser
+ assert walker.data_bridge.kernel_computation is True
+ assert walker.data_bridge.origin_data_source == [{"city": "Paris"}]
+ assert walker.data_bridge.field_specs == [{"fid": "city"}]
+ assert walker.data_bridge.parse_dsl_type == "server"
+ assert walker.data_bridge.dataset_type == "custom_dataset"
+
+
+def test_pygwalker_init_preserves_get_data_parser_override(monkeypatch):
+ monkeypatch.setattr(pygwalker_module, "check_update", lambda: None)
+ calls = []
+
+ class FakeParser:
+ data_size = 1
+ raw_fields = [{"fid": "city"}]
+ dataset_type = "pandas_dataframe"
+ field_metas = []
+
+ def to_records(self, limit=None):
+ return [{"city": "London"}]
+
+ class HookedPygWalker(PygWalker):
+ def _get_data_parser(self, **kwargs):
+ calls.append(kwargs)
+ return FakeParser()
+
+ walker = HookedPygWalker(
+ gid="hooked",
+ dataset=pd.DataFrame([{"city": "London"}]),
+ field_specs=[],
+ spec="",
+ source_invoke_code="",
+ theme_key="g2",
+ appearance="light",
+ show_cloud_tool=False,
+ use_preview=False,
+ kernel_computation=False,
+ cloud_computation=False,
+ use_save_tool=False,
+ is_export_dataframe=False,
+ kanaries_api_key="",
+ default_tab="vis",
+ gw_mode="explore",
+ )
+
+ assert len(calls) == 1
+ assert calls[0]["dataset"].to_dict("records") == [{"city": "London"}]
+ assert walker.data_parser.dataset_type == "pandas_dataframe"
+
+
def test_pygwalker_kernel_callbacks_register_data_query_endpoints(monkeypatch):
walker = _make_walker(monkeypatch, kernel_computation=True)
comm = BaseCommunication("core")
@@ -116,6 +622,116 @@ def test_pygwalker_kernel_data_query_callback_returns_records(monkeypatch):
}
+def test_pygwalker_kernel_data_query_callback_validates_payload(monkeypatch):
+ walker = _make_walker(monkeypatch, kernel_computation=True)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg("get_datas", {})
+
+ assert response["code"] == ErrorCode.INVALID_REQUEST
+ assert "sql" in response["message"]
+
+
+def test_pygwalker_kernel_data_query_callback_rejects_unknown_fields(monkeypatch):
+ walker = _make_walker(monkeypatch, kernel_computation=True)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg("get_datas", {"sql": "SELECT 1", "unexpected": True})
+
+ assert response["code"] == ErrorCode.INVALID_REQUEST
+ assert "unexpected" in response["message"]
+
+
+def test_base_communication_envelope_routes_valid_message(monkeypatch):
+ walker = _make_walker(monkeypatch, kernel_computation=True)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg_envelope(
+ {
+ "action": "get_datas",
+ "data": {"sql": "SELECT SUM(value) AS total FROM pygwalker_mid_table"},
+ "gid": "core",
+ "rid": "request-1",
+ }
+ )
+
+ assert response == {
+ "code": 0,
+ "data": {"datas": [{"total": 3}]},
+ "message": "success",
+ }
+
+
+def test_base_communication_envelope_defaults_missing_data(monkeypatch):
+ walker = _make_walker(monkeypatch)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg_envelope({"action": "ping"})
+
+ assert response == {"code": 0, "data": {}, "message": "success"}
+
+
+def test_ping_callback_rejects_unknown_fields(monkeypatch):
+ walker = _make_walker(monkeypatch)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg("ping", {"unexpected": True})
+
+ assert response["code"] == ErrorCode.INVALID_REQUEST
+ assert "unexpected" in response["message"]
+
+
+def test_base_communication_envelope_rejects_missing_action():
+ comm = BaseCommunication("core")
+
+ response = comm._receive_msg_envelope({"data": {}})
+
+ assert response["code"] == ErrorCode.INVALID_REQUEST
+ assert "action" in response["message"]
+
+
+def test_base_communication_rejects_unknown_action_as_invalid_request():
+ comm = BaseCommunication("core")
+
+ response = comm._receive_msg("missing_endpoint", {})
+
+ assert response == {
+ "code": ErrorCode.INVALID_REQUEST,
+ "data": None,
+ "message": "Unknown action: missing_endpoint",
+ }
+
+
+def test_pygwalker_batch_payload_query_callback_validates_payload(monkeypatch):
+ walker = _make_walker(monkeypatch, kernel_computation=True)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg("batch_get_datas_by_payload", {"queryList": ["not-a-payload"]})
+
+ assert response["code"] == ErrorCode.INVALID_REQUEST
+ assert "queryList" in response["message"] or "query_list" in response["message"]
+
+
+def test_pygwalker_payload_query_callback_rejects_unknown_nested_fields(monkeypatch):
+ walker = _make_walker(monkeypatch, kernel_computation=True)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg(
+ "get_datas_by_payload",
+ {"payload": {"workflow": [{"type": "view"}], "unexpected": True}},
+ )
+
+ assert response["code"] == ErrorCode.INVALID_REQUEST
+ assert "unexpected" in response["message"]
+
+
def test_pygwalker_browser_callbacks_do_not_register_data_query_endpoints(monkeypatch):
walker = _make_walker(monkeypatch, kernel_computation=False)
comm = BaseCommunication("core")
@@ -127,6 +743,17 @@ def test_pygwalker_browser_callbacks_do_not_register_data_query_endpoints(monkey
assert "get_datas_by_payload" not in comm._endpoint_map
+def test_pygwalker_get_latest_vis_spec_callback_rejects_unknown_fields(monkeypatch):
+ walker = _make_walker(monkeypatch, kernel_computation=False)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg("get_latest_vis_spec", {"unexpected": True})
+
+ assert response["code"] == ErrorCode.INVALID_REQUEST
+ assert "unexpected" in response["message"]
+
+
def test_pygwalker_request_data_callback_uploads_current_records(monkeypatch):
upload_calls = []
@@ -137,9 +764,7 @@ def __init__(self, comm):
def run(self, **kwargs):
upload_calls.append(kwargs)
- monkeypatch.setattr(
- pygwalker_module, "BatchUploadDatasToolOnWidgets", FakeUploadTool
- )
+ monkeypatch.setattr(pygwalker_module, "BatchUploadDatasToolOnWidgets", FakeUploadTool)
walker = _make_walker(monkeypatch, kernel_computation=False)
comm = BaseCommunication("core")
walker._init_callback(comm)
@@ -159,29 +784,35 @@ def run(self, **kwargs):
]
+def test_pygwalker_request_data_callback_rejects_unknown_fields(monkeypatch):
+ upload_calls = []
+
+ class FakeUploadTool:
+ def __init__(self, comm):
+ self.comm = comm
+
+ def run(self, **kwargs):
+ upload_calls.append(kwargs)
+
+ monkeypatch.setattr(pygwalker_module, "BatchUploadDatasToolOnWidgets", FakeUploadTool)
+ walker = _make_walker(monkeypatch, kernel_computation=False)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg("request_data", {"unexpected": True})
+
+ assert response["code"] == ErrorCode.INVALID_REQUEST
+ assert "unexpected" in response["message"]
+ assert upload_calls == []
+
+
def test_pygwalker_update_spec_callback_updates_runtime_state(monkeypatch):
walker = _make_walker(monkeypatch, use_save_tool=True)
comm = BaseCommunication("core")
walker._init_callback(comm)
vis_spec = [{"name": "Updated chart", "encodings": {}}]
workflow_list = [{"type": "filter"}]
- chart_data = {
- "charts": [
- {
- "rowIndex": 0,
- "colIndex": 0,
- "data": "data:image/png;base64,abc",
- "height": 100,
- "width": 200,
- "canvasHeight": 100,
- "canvasWidth": 200,
- }
- ],
- "singleChart": "data:image/png;base64,abc",
- "nRows": 1,
- "nCols": 1,
- "title": "Updated chart",
- }
+ chart_data = _chart_payload()
response = comm._receive_msg(
"update_spec",
@@ -192,12 +823,172 @@ def test_pygwalker_update_spec_callback_updates_runtime_state(monkeypatch):
},
)
- assert response == {"code": 0, "data": None, "message": "success"}
+ assert response == {"code": 0, "data": {}, "message": "success"}
assert walker.vis_spec == vis_spec
assert walker.workflow_list == workflow_list
assert walker._chart_map["Updated chart"].title == "Updated chart"
+def test_pygwalker_update_spec_callback_validates_payload(monkeypatch):
+ walker = _make_walker(monkeypatch, use_save_tool=True)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg("update_spec", {"visSpec": []})
+
+ assert response["code"] == ErrorCode.INVALID_REQUEST
+ assert "chartData" in response["message"] or "chart_data" in response["message"]
+
+
+def test_pygwalker_update_spec_callback_validates_chart_data_shape(monkeypatch):
+ walker = _make_walker(monkeypatch, use_save_tool=True)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg(
+ "update_spec",
+ {
+ "visSpec": [{"name": "Broken chart", "encodings": {}}],
+ "chartData": {"title": "Broken chart"},
+ },
+ )
+
+ assert response["code"] == ErrorCode.INVALID_REQUEST
+ assert "singleChart" in response["message"] or "single_chart" in response["message"]
+
+
+def test_pygwalker_update_spec_callback_rejects_extra_chart_data_fields(monkeypatch):
+ walker = _make_walker(monkeypatch, use_save_tool=True)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+ chart_data = _chart_payload("Extra chart")
+ chart_data["mode"] = "explore"
+
+ response = comm._receive_msg(
+ "update_spec",
+ {
+ "visSpec": [{"name": "Extra chart", "encodings": {}}],
+ "chartData": chart_data,
+ },
+ )
+
+ assert response["code"] == ErrorCode.INVALID_REQUEST
+ assert "mode" in response["message"]
+
+
+def test_pygwalker_update_spec_callback_defaults_missing_workflow_list(monkeypatch):
+ walker = _make_walker(monkeypatch, use_save_tool=True)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+ vis_spec = [{"name": "Default workflow chart", "encodings": {}}]
+
+ response = comm._receive_msg(
+ "update_spec",
+ {
+ "visSpec": vis_spec,
+ "chartData": _chart_payload("Default workflow chart"),
+ },
+ )
+
+ assert response == {"code": 0, "data": {}, "message": "success"}
+ assert walker.vis_spec == vis_spec
+ assert walker.workflow_list == []
+
+
+def test_pygwalker_runtime_spec_properties_remain_writable(monkeypatch):
+ walker = _make_walker(monkeypatch)
+ vis_spec = [{"name": "Manual chart", "encodings": {}}]
+
+ walker.spec_type = "manual"
+ walker.spec_version = "0.6.0"
+ walker.vis_spec = vis_spec
+ walker.workflow_list = [{"workflow": []}]
+ walker._chart_map = {}
+ walker._chart_name_index_map = {"Manual chart": 0}
+
+ assert walker.spec_type == "manual"
+ assert walker.spec_version == "0.6.0"
+ assert walker.vis_spec == vis_spec
+ assert walker.workflow_list == [{"workflow": []}]
+ assert walker._chart_map == {}
+ assert walker._chart_name_index_map == {"Manual chart": 0}
+
+
+def test_pygwalker_update_spec_callback_writes_json_file(monkeypatch, tmp_path):
+ spec_path = tmp_path / "gw_config.json"
+ spec_path.write_text(json.dumps({"config": [], "chart_map": {}, "workflow_list": [], "version": "0.5.0"}))
+ walker = _make_walker(monkeypatch, spec=str(spec_path), use_save_tool=True)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+ vis_spec = [{"name": "File chart", "encodings": {}}]
+ workflow_list = [{"workflow": [{"type": "view"}]}]
+
+ response = comm._receive_msg(
+ "update_spec",
+ {
+ "visSpec": vis_spec,
+ "workflowList": workflow_list,
+ "chartData": _chart_payload("File chart"),
+ },
+ )
+
+ assert response == {"code": 0, "data": {}, "message": "success"}
+ assert json.loads(spec_path.read_text()) == {
+ "config": vis_spec,
+ "chart_map": {},
+ "version": __version__,
+ "workflow_list": workflow_list,
+ }
+
+
+def test_pygwalker_loads_saved_chart_map_from_spec(monkeypatch):
+ chart_payload = _chart_payload("Saved chart")
+
+ walker = _make_walker(
+ monkeypatch,
+ spec={
+ "config": [],
+ "chart_map": {"Saved chart": chart_payload},
+ "workflow_list": [],
+ "version": "0.5.0",
+ },
+ )
+
+ assert walker.chart_list == ["Saved chart"]
+ assert walker._get_chart_by_name("Saved chart").title == "Saved chart"
+
+
+def test_pygwalker_to_code_exports_current_spec_state(monkeypatch):
+ walker = _make_walker(monkeypatch)
+ walker.vis_spec = [{"name": "Bob's chart", "encodings": {}}]
+ walker.workflow_list = [{"workflow": [{"type": "view"}]}]
+ walker.spec_version = "0.6.0"
+
+ code = walker.to_code(dataset_name="source_df", variable_name="explorer")
+
+ assert code.startswith("import pygwalker as pyg\n\n")
+ assert code.endswith("explorer = pyg.walk(source_df, spec=spec)")
+
+ namespace = {}
+ exec(code.splitlines()[2], {}, namespace)
+ assert json.loads(namespace["spec"]) == {
+ "config": [{"name": "Bob's chart", "encodings": {}}],
+ "chart_map": {},
+ "version": "0.6.0",
+ "workflow_list": [{"workflow": [{"type": "view"}]}],
+ }
+
+
+def test_pygwalker_to_code_can_omit_import_and_rejects_invalid_variable_name(monkeypatch):
+ walker = _make_walker(monkeypatch)
+
+ assert walker.to_code(include_import=False).startswith("spec = ")
+ with pytest.raises(ValueError, match="variable_name"):
+ walker.to_code(variable_name="not valid")
+ with pytest.raises(ValueError, match="variable_name"):
+ walker.to_code(variable_name="class")
+
+
def test_pygwalker_export_dataframe_callback_stores_last_dataframe(monkeypatch):
previous_exported_dataframe = GlobalVarManager.last_exported_dataframe
walker = _make_walker(monkeypatch, is_export_dataframe=True)
@@ -210,23 +1001,311 @@ def test_pygwalker_export_dataframe_callback_stores_last_dataframe(monkeypatch):
{"sql": "SELECT city, value FROM pygwalker_mid_table WHERE value = 2"},
)
- assert response == {"code": 0, "data": None, "message": "success"}
- assert walker.last_exported_dataframe.to_dict("records") == [
- {"city": "Tokyo", "value": 2}
- ]
- assert (
- GlobalVarManager.last_exported_dataframe is walker.last_exported_dataframe
+ assert response == {"code": 0, "data": {}, "message": "success"}
+ assert walker.last_exported_dataframe.to_dict("records") == [{"city": "Tokyo", "value": 2}]
+ assert GlobalVarManager.last_exported_dataframe is walker.last_exported_dataframe
+ finally:
+ GlobalVarManager.last_exported_dataframe = previous_exported_dataframe
+
+
+def test_pygwalker_export_dataframe_by_payload_callback_stores_last_dataframe(monkeypatch):
+ previous_exported_dataframe = GlobalVarManager.last_exported_dataframe
+ walker = _make_walker(monkeypatch, is_export_dataframe=True)
+ calls = []
+ walker.data_parser = SimpleNamespace(
+ get_datas_by_payload=lambda payload: calls.append(payload) or [{"city": "London", "value": 1}]
+ )
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ try:
+ response = comm._receive_msg(
+ "export_dataframe_by_payload",
+ {"payload": {"workflow": [{"type": "view"}]}},
)
+
+ assert response == {"code": 0, "data": {}, "message": "success"}
+ assert calls == [{"workflow": [{"type": "view"}]}]
+ assert walker.last_exported_dataframe.to_dict("records") == [{"city": "London", "value": 1}]
+ assert GlobalVarManager.last_exported_dataframe is walker.last_exported_dataframe
finally:
GlobalVarManager.last_exported_dataframe = previous_exported_dataframe
+def test_pygwalker_batch_sql_callback_returns_records(monkeypatch):
+ walker = _make_walker(monkeypatch, kernel_computation=True)
+ calls = []
+ walker.data_parser = SimpleNamespace(
+ batch_get_datas_by_sql=lambda query_list: calls.append(query_list) or [[{"total": 3}]]
+ )
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg("batch_get_datas_by_sql", {"queryList": ["SELECT 1"]})
+
+ assert response == {"code": 0, "data": {"datas": [[{"total": 3}]]}, "message": "success"}
+ assert calls == [["SELECT 1"]]
+
+
+def test_pygwalker_upload_spec_to_cloud_callback_writes_workspace_path(monkeypatch):
+ walker = _make_walker(monkeypatch, use_save_tool=True)
+ writes = []
+ walker.cloud_service = SimpleNamespace(
+ get_kanaries_user_info=lambda: {"workspaceName": "workspace"},
+ write_config_to_cloud=lambda path, data: writes.append((path, json.loads(data))),
+ )
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg("upload_spec_to_cloud", {"newToken": "", "fileName": "chart.json"})
+
+ assert response == {
+ "code": 0,
+ "data": {"specFilePath": "workspace/chart.json"},
+ "message": "success",
+ }
+ assert writes == [
+ (
+ "workspace/chart.json",
+ {
+ "config": [],
+ "chart_map": {},
+ "workflow_list": [],
+ "version": __version__,
+ },
+ )
+ ]
+
+
+def test_pygwalker_upload_spec_to_cloud_callback_validates_payload(monkeypatch):
+ walker = _make_walker(monkeypatch, use_save_tool=True)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg("upload_spec_to_cloud", {"newToken": ""})
+
+ assert response["code"] == ErrorCode.INVALID_REQUEST
+ assert "fileName" in response["message"] or "file_name" in response["message"]
+
+
+def test_pygwalker_save_chart_callback_validates_and_stores_chart(monkeypatch):
+ walker = _make_walker(monkeypatch, use_save_tool=True)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg("save_chart", _chart_payload("Recovered chart"))
+
+ assert response == {"code": 0, "data": {}, "message": "success"}
+ assert walker.chart_list == ["Recovered chart"]
+ assert walker._get_chart_by_name("Recovered chart").single_chart == "data:image/png;base64,abc"
+
+
+def test_pygwalker_save_chart_callback_validates_payload(monkeypatch):
+ walker = _make_walker(monkeypatch, use_save_tool=True)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg("save_chart", {"title": "Broken chart"})
+
+ assert response["code"] == ErrorCode.INVALID_REQUEST
+ assert "singleChart" in response["message"] or "single_chart" in response["message"]
+
+
+def test_pygwalker_cloud_text_callbacks_validate_payloads(monkeypatch):
+ ask_calls = []
+ chat_calls = []
+ walker = _make_walker(
+ monkeypatch,
+ show_cloud_tool=True,
+ custom_ask_callback=lambda metas, query: ask_calls.append((metas, query)) or {"chart": "bar"},
+ custom_chat_callback=lambda metas, chats: chat_calls.append((metas, chats)) or {"chart": "line"},
+ )
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ ask_response = comm._receive_msg("get_spec_by_text", {"metas": [{"fid": "city"}], "query": "show city"})
+ chat_response = comm._receive_msg(
+ "get_chart_by_chats",
+ {"metas": [{"fid": "city"}], "chats": [{"role": "user", "content": "show city"}]},
+ )
+ invalid_response = comm._receive_msg("get_spec_by_text", {"metas": []})
+
+ assert ask_response == {"code": 0, "data": {"data": {"chart": "bar"}}, "message": "success"}
+ assert chat_response == {"code": 0, "data": {"data": {"chart": "line"}}, "message": "success"}
+ assert ask_calls == [([{"fid": "city"}], "show city")]
+ assert chat_calls == [([{"fid": "city"}], [{"role": "user", "content": "show city"}])]
+ assert invalid_response["code"] == ErrorCode.INVALID_REQUEST
+ assert "query" in invalid_response["message"]
+
+
+def test_pygwalker_upload_cloud_chart_callback_validates_and_uploads(monkeypatch):
+ upload_calls = []
+ walker = _make_walker(monkeypatch, show_cloud_tool=True)
+ walker.cloud_service = SimpleNamespace(
+ upload_cloud_chart=lambda **kwargs: (
+ upload_calls.append(kwargs) or {"chart_id": "chart-id", "dataset_id": "dataset-id"}
+ )
+ )
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg(
+ "upload_to_cloud_charts",
+ {
+ "chartName": "Cloud chart",
+ "datasetName": "Dataset",
+ "isPublic": True,
+ "visSpec": [{"name": "Cloud chart"}],
+ "workflow": [{"type": "view"}],
+ },
+ )
+ invalid_response = comm._receive_msg(
+ "upload_to_cloud_charts",
+ {
+ "chartName": "Cloud chart",
+ "datasetName": "Dataset",
+ "isPublic": True,
+ "visSpec": [{"name": "Cloud chart"}],
+ },
+ )
+
+ assert response == {
+ "code": 0,
+ "data": {"chartId": "chart-id", "datasetId": "dataset-id"},
+ "message": "success",
+ }
+ assert upload_calls == [
+ {
+ "data_parser": walker.data_parser,
+ "chart_name": "Cloud chart",
+ "dataset_name": "Dataset",
+ "workflow": [{"type": "view"}],
+ "spec_list": [{"name": "Cloud chart"}],
+ "is_public": True,
+ }
+ ]
+ assert invalid_response["code"] == ErrorCode.INVALID_REQUEST
+ assert "workflow" in invalid_response["message"]
+
+
+def test_pygwalker_upload_cloud_dashboard_callback_validates_and_uploads(monkeypatch):
+ upload_calls = []
+ walker = _make_walker(monkeypatch, show_cloud_tool=True)
+ walker.cloud_service = SimpleNamespace(
+ upload_cloud_dashboard=lambda **kwargs: (
+ upload_calls.append(kwargs) or {"dashboard_id": "dashboard-id", "dataset_id": "dataset-id"}
+ )
+ )
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg(
+ "upload_to_cloud_dashboard",
+ {
+ "chartName": "Dashboard",
+ "datasetName": "Dataset",
+ "isPublic": False,
+ "isCreateDashboard": True,
+ "visSpec": [{"name": "Chart A"}],
+ "workflowList": [[{"type": "view"}]],
+ },
+ )
+ invalid_response = comm._receive_msg(
+ "upload_to_cloud_dashboard",
+ {
+ "chartName": "Dashboard",
+ "datasetName": "Dataset",
+ "isPublic": False,
+ "visSpec": [{"name": "Chart A"}],
+ "workflowList": [[{"type": "view"}]],
+ },
+ )
+
+ assert response == {
+ "code": 0,
+ "data": {"dashboardId": "dashboard-id", "datasetId": "dataset-id"},
+ "message": "success",
+ }
+ assert upload_calls == [
+ {
+ "data_parser": walker.data_parser,
+ "dashboard_name": "Dashboard",
+ "dataset_name": "Dataset",
+ "workflow_list": [[{"type": "view"}]],
+ "spec_list": [{"name": "Chart A"}],
+ "is_public": False,
+ "create_dashboard_flag": True,
+ "appearance": walker.appearance,
+ }
+ ]
+ assert invalid_response["code"] == ErrorCode.INVALID_REQUEST
+ assert "isCreateDashboard" in invalid_response["message"] or "is_create_dashboard" in invalid_response["message"]
+
+
+def test_pygwalker_open_in_desktop_callback_encodes_payload(monkeypatch):
+ links = []
+ monkeypatch.setattr(
+ desktop_import_module.DesktopImportService,
+ "_open_platform_link",
+ staticmethod(lambda link: links.append(link)),
+ )
+ walker = _make_walker(monkeypatch)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg(
+ "open_in_desktop",
+ {
+ "spec": [{"name": "Chart"}],
+ "fields": [{"fid": "city"}],
+ },
+ )
+
+ assert response == {"code": 0, "data": {}, "message": "success"}
+ parsed = urllib.parse.urlparse(links[0])
+ query = urllib.parse.parse_qs(parsed.query)
+
+ def _decode_query_value(name):
+ compressed = base64.b64decode(urllib.parse.unquote(query[name][0]))
+ return json.loads(zlib.decompress(compressed).decode())
+
+ assert parsed.scheme == "gw"
+ assert parsed.netloc == "import"
+ assert _decode_query_value("spec") == [{"name": "Chart"}]
+ assert _decode_query_value("fields") == [{"fid": "city"}]
+ assert _decode_query_value("data") == [
+ {"city": "London", "value": 1},
+ {"city": "Tokyo", "value": 2},
+ ]
+
+
+def test_pygwalker_open_in_desktop_callback_validates_payload(monkeypatch):
+ links = []
+ monkeypatch.setattr(
+ desktop_import_module.DesktopImportService,
+ "_open_platform_link",
+ staticmethod(lambda link: links.append(link)),
+ )
+ walker = _make_walker(monkeypatch)
+ comm = BaseCommunication("core")
+ walker._init_callback(comm)
+
+ response = comm._receive_msg("open_in_desktop", {"spec": []})
+
+ assert response["code"] == ErrorCode.INVALID_REQUEST
+ assert "fields" in response["message"]
+ assert links == []
+
+
@pytest.mark.parametrize(
("kwargs", "expected_kernel_computation"),
[
({}, False),
({"kernel_computation": True}, True),
- ({"env": "JupyterConvert", "kernel_computation": True}, False),
+ ({"computation": "browser"}, False),
+ ({"computation": "kernel"}, True),
+ ({"computation": "cloud"}, False),
+ ({"env": "JupyterConvert", "kernel_computation": False}, False),
],
)
def test_jupyter_walk_sets_pygwalker_kernel_computation_mode(
@@ -239,17 +1318,385 @@ def test_jupyter_walk_sets_pygwalker_kernel_computation_mode(
monkeypatch.setattr(jupyter, "check_kaggle", lambda: False)
monkeypatch.setattr(jupyter, "check_convert", lambda: False)
monkeypatch.setattr(jupyter, "get_kaggle_run_type", lambda: "")
+ monkeypatch.setattr(PygWalker, "display_on_jupyter_use_anywidget", lambda self: None)
monkeypatch.setattr(PygWalker, "display_on_jupyter_use_widgets", lambda self: None)
monkeypatch.setattr(PygWalker, "display_on_jupyter", lambda self: None)
monkeypatch.setattr(PygWalker, "display_on_convert_html", lambda self: None)
+ cloud_uploads = _patch_cloud_computation_parser(monkeypatch) if kwargs.get("computation") == "cloud" else []
+
+ with _expected_legacy_computation_warning(kwargs):
+ walker = jupyter.walk(
+ pd.DataFrame([{"city": "London", "value": 1}]),
+ gid="entry",
+ **kwargs,
+ )
+
+ assert walker.kernel_computation is expected_kernel_computation
+ if kwargs.get("computation") == "cloud":
+ assert walker.cloud_computation is True
+ assert walker.dataset_type == "cloud_dataset"
+ assert len(cloud_uploads) == 1
+
+
+@pytest.mark.parametrize(
+ "kwargs",
+ [
+ {"computation": "kernel"},
+ {"computation": "cloud"},
+ {"kernel_computation": True},
+ {"use_kernel_calc": True},
+ {"cloud_computation": True},
+ ],
+)
+def test_jupyter_walk_rejects_live_computation_for_convert_env(monkeypatch, kwargs):
+ monkeypatch.setattr(jupyter, "check_kaggle", lambda: False)
+ monkeypatch.setattr(jupyter, "check_convert", lambda: False)
+ monkeypatch.setattr(jupyter, "get_kaggle_run_type", lambda: "")
+
+ with _expected_legacy_computation_warning(kwargs):
+ with pytest.raises(ValueError, match="JupyterConvert/static HTML output does not support"):
+ jupyter.walk(
+ pd.DataFrame([{"city": "London", "value": 1}]),
+ env="JupyterConvert",
+ **kwargs,
+ )
+
+
+def test_jupyter_walk_accepts_explicit_spec_path(monkeypatch, tmp_path):
+ monkeypatch.setattr(pygwalker_module, "check_update", lambda: None)
+ monkeypatch.setattr(pygwalker_module, "track_event", lambda *_args, **_kwargs: None)
+ monkeypatch.setattr(jupyter, "check_kaggle", lambda: False)
+ monkeypatch.setattr(jupyter, "check_convert", lambda: False)
+ monkeypatch.setattr(jupyter, "get_kaggle_run_type", lambda: "")
+ monkeypatch.setattr(PygWalker, "display_on_jupyter_use_anywidget", lambda self: None)
+
+ spec_path = tmp_path / "gw_config.json"
+ spec_path.write_text(json.dumps({"config": [], "chart_map": {}, "workflow_list": [], "version": "0.5.0"}))
walker = jupyter.walk(
pd.DataFrame([{"city": "London", "value": 1}]),
- gid="entry",
- **kwargs,
+ gid="spec-path",
+ spec_path=str(spec_path),
)
- assert walker.kernel_computation is expected_kernel_computation
+ assert walker.spec_manager.spec == str(spec_path)
+ assert walker.spec_manager.spec_type == "json_file"
+
+
+def test_jupyter_walk_rejects_spec_and_spec_path(monkeypatch, tmp_path):
+ monkeypatch.setattr(jupyter, "check_kaggle", lambda: False)
+ monkeypatch.setattr(jupyter, "check_convert", lambda: False)
+ monkeypatch.setattr(jupyter, "get_kaggle_run_type", lambda: "")
+
+ with pytest.raises(ValueError, match="Pass only one of `spec` or `spec_path`"):
+ jupyter.walk(
+ pd.DataFrame([{"city": "London", "value": 1}]),
+ spec="{}",
+ spec_path=str(tmp_path / "gw_config.json"),
+ )
+
+
+def test_jupyter_walk_accepts_public_walker_object(monkeypatch):
+ import pygwalker
+
+ monkeypatch.setattr(pygwalker_module, "check_update", lambda: None)
+ monkeypatch.setattr(pygwalker_module, "track_event", lambda *_args, **_kwargs: None)
+ monkeypatch.setattr(jupyter, "check_kaggle", lambda: False)
+ monkeypatch.setattr(jupyter, "check_convert", lambda: False)
+ monkeypatch.setattr(jupyter, "get_kaggle_run_type", lambda: "")
+
+ display_calls = []
+ monkeypatch.setattr(
+ PygWalker,
+ "display_on_jupyter_use_anywidget",
+ lambda self: display_calls.append(self.gid),
+ )
+
+ public_walker = pygwalker.Walker(
+ pd.DataFrame([{"city": "London", "value": 1}]),
+ gid="public-jupyter",
+ computation="browser",
+ )
+ result = jupyter.walk(public_walker)
+
+ assert result is public_walker.core
+ assert display_calls == ["public-jupyter"]
+
+
+def test_public_walker_accepts_empty_dataframe(monkeypatch):
+ import pygwalker
+
+ monkeypatch.setattr(pygwalker_module, "check_update", lambda: None)
+ monkeypatch.setattr(pygwalker_module, "track_event", lambda *_args, **_kwargs: None)
+
+ public_walker = pygwalker.Walker(
+ pd.DataFrame({"city": pd.Series(dtype="object"), "value": pd.Series(dtype="int64")}),
+ gid="public-empty",
+ computation="browser",
+ )
+
+ assert public_walker.core.gid == "public-empty"
+ assert public_walker.core.origin_data_source == []
+ assert public_walker.core.dataset_type == "pandas_dataframe"
+ assert public_walker.core.parse_dsl_type == "client"
+ assert [field["fid"] for field in public_walker.core.field_specs] == ["city", "value"]
+
+
+def test_public_walker_accepts_pyarrow_table(monkeypatch):
+ import pygwalker
+
+ monkeypatch.setattr(pygwalker_module, "check_update", lambda: None)
+ monkeypatch.setattr(pygwalker_module, "track_event", lambda *_args, **_kwargs: None)
+
+ public_walker = pygwalker.Walker(
+ pa.table({"city": ["London"], "value": [1]}),
+ gid="public-pyarrow",
+ computation="browser",
+ )
+
+ assert public_walker.core.gid == "public-pyarrow"
+ assert public_walker.core.origin_data_source == [{"city": "London", "value": 1}]
+ assert public_walker.core.dataset_type == "pyarrow_table"
+ assert public_walker.core.parse_dsl_type == "client"
+ assert [field["fid"] for field in public_walker.core.field_specs] == ["city", "value"]
+
+
+def test_jupyter_walk_public_walker_legacy_widget_env_warns_once(monkeypatch):
+ from pygwalker.api.walker import Walker
+
+ monkeypatch.setattr(pygwalker_module, "check_update", lambda: None)
+ monkeypatch.setattr(pygwalker_module, "track_event", lambda *_args, **_kwargs: None)
+ monkeypatch.setattr(jupyter, "check_kaggle", lambda: False)
+ monkeypatch.setattr(jupyter, "check_convert", lambda: False)
+ monkeypatch.setattr(jupyter, "get_kaggle_run_type", lambda: "")
+
+ display_calls = []
+ monkeypatch.setattr(
+ PygWalker,
+ "display_on_jupyter_use_widgets",
+ lambda self, iframe_width=None, iframe_height=None: display_calls.append(
+ (self.gid, iframe_width, iframe_height)
+ ),
+ )
+
+ public_walker = Walker(
+ pd.DataFrame([{"city": "London", "value": 1}]),
+ gid="public-legacy-widget",
+ computation="browser",
+ )
+
+ with pytest.warns(DeprecationWarning, match="legacy Jupyter transport") as warnings:
+ result = jupyter.walk(public_walker, env="JupyterWidget")
+
+ assert len(warnings) == 1
+ assert result is public_walker.core
+ assert display_calls == [("public-legacy-widget", None, None)]
+
+
+def test_walker_show_legacy_inline_env_warns_once_with_core_display(monkeypatch):
+ from pygwalker.api.walker import Walker
+
+ monkeypatch.setattr(pygwalker_module, "check_update", lambda: None)
+ monkeypatch.setattr(pygwalker_module, "track_event", lambda *_args, **_kwargs: None)
+
+ display_calls = []
+ public_walker = Walker(
+ pd.DataFrame([{"city": "London", "value": 1}]),
+ gid="public-legacy-inline",
+ computation="browser",
+ )
+ public_walker.core.jupyter_display_manager = SimpleNamespace(
+ display_on_jupyter=lambda: display_calls.append(public_walker.core.gid)
+ )
+
+ with pytest.warns(DeprecationWarning, match="legacy Jupyter transport") as warnings:
+ result = public_walker.show("Jupyter")
+
+ assert len(warnings) == 1
+ assert result is public_walker
+ assert display_calls == ["public-legacy-inline"]
+
+
+def test_jupyter_walk_legacy_widget_env_uses_ipywidgets_transport(monkeypatch):
+ monkeypatch.setattr(pygwalker_module, "check_update", lambda: None)
+ monkeypatch.setattr(pygwalker_module, "track_event", lambda *_args, **_kwargs: None)
+ monkeypatch.setattr(jupyter, "check_kaggle", lambda: False)
+ monkeypatch.setattr(jupyter, "check_convert", lambda: False)
+ monkeypatch.setattr(jupyter, "get_kaggle_run_type", lambda: "")
+
+ display_calls = []
+ monkeypatch.setattr(
+ PygWalker,
+ "display_on_jupyter_use_widgets",
+ lambda self, iframe_width=None, iframe_height=None: display_calls.append(
+ (self.gid, iframe_width, iframe_height)
+ ),
+ )
+
+ with pytest.warns(DeprecationWarning, match="legacy Jupyter transport"):
+ walker = jupyter.walk(pd.DataFrame([{"city": "London", "value": 1}]), gid="legacy-widget", env="JupyterWidget")
+
+ assert walker.gid == "legacy-widget"
+ assert display_calls == [("legacy-widget", None, None)]
+
+
+def test_jupyter_walk_legacy_inline_env_warns_and_uses_inline_transport(monkeypatch):
+ monkeypatch.setattr(pygwalker_module, "check_update", lambda: None)
+ monkeypatch.setattr(pygwalker_module, "track_event", lambda *_args, **_kwargs: None)
+ monkeypatch.setattr(jupyter, "check_kaggle", lambda: False)
+ monkeypatch.setattr(jupyter, "check_convert", lambda: False)
+ monkeypatch.setattr(jupyter, "get_kaggle_run_type", lambda: "")
+
+ display_calls = []
+ monkeypatch.setattr(
+ jupyter_display_module.JupyterDisplayManager,
+ "display_on_jupyter",
+ lambda self: display_calls.append(self.walker.gid),
+ )
+
+ with pytest.warns(DeprecationWarning, match="legacy Jupyter transport") as warnings:
+ walker = jupyter.walk(pd.DataFrame([{"city": "London", "value": 1}]), gid="legacy-inline", env="Jupyter")
+
+ assert len(warnings) == 1
+ assert walker.gid == "legacy-inline"
+ assert display_calls == ["legacy-inline"]
+
+
+def test_core_legacy_jupyter_display_methods_warn(monkeypatch):
+ inline_calls = []
+ widget_calls = []
+ walker = _make_walker(monkeypatch, gid="core-legacy")
+ walker.jupyter_display_manager = SimpleNamespace(
+ display_on_jupyter=lambda: inline_calls.append(walker.gid),
+ display_on_jupyter_use_widgets=lambda iframe_width=None, iframe_height=None: widget_calls.append(
+ (iframe_width, iframe_height)
+ ),
+ )
+
+ with pytest.warns(DeprecationWarning, match="legacy Jupyter transport"):
+ walker.display_on_jupyter()
+ with pytest.warns(DeprecationWarning, match="legacy Jupyter transport"):
+ walker.display_on_jupyter_use_widgets("640px", "480px")
+
+ assert inline_calls == ["core-legacy"]
+ assert widget_calls == [("640px", "480px")]
+
+
+def test_display_on_jupyter_anywidget_sends_browser_data(monkeypatch):
+ from pygwalker.services import anywidget_widget
+
+ displayed = []
+ created = []
+ monkeypatch.setattr(pygwalker_module, "display_html", lambda widget: displayed.append(widget))
+ monkeypatch.setattr(
+ anywidget_widget,
+ "create_anywidget_for_walker",
+ lambda walker, *, env, data_source: created.append((walker, env, data_source)) or "widget",
+ )
+
+ walker = _make_walker(monkeypatch, kernel_computation=False, cloud_computation=False, use_preview=True)
+ walker.display_on_jupyter_use_anywidget()
+
+ assert walker.use_preview is False
+ assert created == [(walker, "anywidget", walker.origin_data_source)]
+ assert displayed == ["widget"]
+
+
+def test_display_on_jupyter_anywidget_uses_comm_data_for_live_computation(monkeypatch):
+ from pygwalker.services import anywidget_widget
+
+ displayed = []
+ created = []
+ monkeypatch.setattr(pygwalker_module, "display_html", lambda widget: displayed.append(widget))
+ monkeypatch.setattr(
+ anywidget_widget,
+ "create_anywidget_for_walker",
+ lambda walker, *, env, data_source: created.append((walker, env, data_source)) or "widget",
+ )
+
+ walker = _make_walker(monkeypatch, kernel_computation=True, cloud_computation=False, use_preview=True)
+ walker.display_on_jupyter_use_anywidget()
+
+ assert walker.use_preview is False
+ assert created == [(walker, "anywidget", [])]
+ assert displayed == ["widget"]
+
+
+def test_display_on_jupyter_anywidget_sends_cloud_mode_data(monkeypatch):
+ from pygwalker.services import anywidget_widget
+
+ displayed = []
+ created = []
+ cloud_uploads = _patch_cloud_computation_parser(monkeypatch)
+ monkeypatch.setattr(pygwalker_module, "display_html", lambda widget: displayed.append(widget))
+ monkeypatch.setattr(
+ anywidget_widget,
+ "create_anywidget_for_walker",
+ lambda walker, *, env, data_source: created.append((walker, env, data_source)) or "widget",
+ )
+
+ walker = _make_walker(monkeypatch, kernel_computation=False, cloud_computation=True, use_preview=True)
+ walker.display_on_jupyter_use_anywidget()
+
+ assert walker.use_preview is False
+ assert walker.dataset_type == "cloud_dataset"
+ assert len(cloud_uploads) == 1
+ assert created == [(walker, "anywidget", walker.origin_data_source)]
+ assert displayed == ["widget"]
+
+
+def test_jupyter_walk_public_walker_rejects_rebuilding_params(monkeypatch, tmp_path):
+ from pygwalker.api.walker import Walker
+
+ monkeypatch.setattr(pygwalker_module, "check_update", lambda: None)
+ monkeypatch.setattr(pygwalker_module, "track_event", lambda *_args, **_kwargs: None)
+ monkeypatch.setattr(jupyter, "check_kaggle", lambda: False)
+ monkeypatch.setattr(jupyter, "check_convert", lambda: False)
+ monkeypatch.setattr(jupyter, "get_kaggle_run_type", lambda: "")
+
+ public_walker = Walker(pd.DataFrame([{"city": "London", "value": 1}]), computation="browser")
+
+ with pytest.raises(ValueError, match="cannot apply construction parameters: spec_path"):
+ jupyter.walk(public_walker, spec_path=str(tmp_path / "other.json"))
+
+
+def test_jupyter_walk_public_walker_uses_convert_guard(monkeypatch):
+ from pygwalker.api.walker import Walker
+
+ monkeypatch.setattr(pygwalker_module, "check_update", lambda: None)
+ monkeypatch.setattr(pygwalker_module, "track_event", lambda *_args, **_kwargs: None)
+ monkeypatch.setattr(jupyter, "check_kaggle", lambda: False)
+ monkeypatch.setattr(jupyter, "check_convert", lambda: True)
+ monkeypatch.setattr(jupyter, "get_kaggle_run_type", lambda: "")
+
+ public_walker = Walker(pd.DataFrame([{"city": "London", "value": 1}]), computation="kernel")
+
+ with pytest.raises(ValueError, match="JupyterConvert/static HTML output does not support kernel computation"):
+ jupyter.walk(public_walker)
+
+
+def test_jupyter_walk_public_walker_uses_preview_env(monkeypatch):
+ from pygwalker.api.walker import Walker
+
+ monkeypatch.setattr(pygwalker_module, "check_update", lambda: None)
+ monkeypatch.setattr(pygwalker_module, "track_event", lambda *_args, **_kwargs: None)
+ monkeypatch.setattr(jupyter, "check_kaggle", lambda: False)
+ monkeypatch.setattr(jupyter, "check_convert", lambda: False)
+ monkeypatch.setattr(jupyter, "get_kaggle_run_type", lambda: "batch")
+ monkeypatch.setattr(jupyter, "adjust_kaggle_default_font_size", lambda: None)
+
+ display_calls = []
+ monkeypatch.setattr(PygWalker, "display_preview_on_jupyter", lambda self: display_calls.append(self.gid))
+
+ public_walker = Walker(
+ pd.DataFrame([{"city": "London", "value": 1}]),
+ gid="public-preview",
+ computation="browser",
+ )
+ result = jupyter.walk(public_walker)
+
+ assert result is public_walker.core
+ assert display_calls == ["public-preview"]
def test_to_html_returns_iframe_for_pygwalker_static_export(monkeypatch):
@@ -269,6 +1716,82 @@ def test_to_html_returns_iframe_for_pygwalker_static_export(monkeypatch):
assert 'width="640px"' in rendered
assert 'height="480px"' in rendered
assert "srcdoc=" in rendered
+ assert "eval(script)" not in rendered
+ assert "URL.createObjectURL" in rendered
+
+
+def test_to_html_accepts_pyarrow_table(monkeypatch):
+ monkeypatch.setattr(pygwalker_module, "check_update", lambda: None)
+ monkeypatch.setattr(pygwalker_module, "track_event", lambda *_args, **_kwargs: None)
+ monkeypatch.setattr(pygwalker_module, "get_local_user_id", lambda: "test-user")
+
+ rendered = html.to_html(
+ pa.table({"city": ["London"], "value": [1]}),
+ gid="static-pyarrow",
+ computation="browser",
+ )
+
+ assert 'id="gwalker-static-pyarrow"' in rendered
+ assert "srcdoc=" in rendered
+
+
+@pytest.mark.parametrize(
+ "kwargs",
+ [
+ {"kernel_computation": True},
+ {"cloud_computation": True},
+ {"use_kernel_calc": True},
+ {"computation": "kernel"},
+ {"computation": "cloud"},
+ ],
+)
+def test_to_html_rejects_live_computation_modes(kwargs):
+ with pytest.raises(ValueError, match="Static HTML export does not support kernel or cloud computation"):
+ html.to_html(pd.DataFrame([{"city": "London", "value": 1}]), **kwargs)
+
+
+def test_to_html_rejects_invalid_computation_mode():
+ with pytest.raises(ValueError, match="`computation` must be one of"):
+ html.to_html(pd.DataFrame([{"city": "London", "value": 1}]), computation="server")
+
+
+def test_to_html_allows_disabled_computation_kwargs(monkeypatch):
+ monkeypatch.setattr(pygwalker_module, "check_update", lambda: None)
+ monkeypatch.setattr(pygwalker_module, "track_event", lambda *_args, **_kwargs: None)
+ monkeypatch.setattr(pygwalker_module, "get_local_user_id", lambda: "test-user")
+
+ rendered = html.to_html(
+ pd.DataFrame([{"city": "London", "value": 1}]),
+ kernel_computation=False,
+ cloud_computation=False,
+ use_kernel_calc=None,
+ computation="browser",
+ )
+
+ assert 'id="gwalker-' in rendered
+ assert "srcdoc=" in rendered
+
+
+@pytest.mark.parametrize("render_type", ["explorer", "profiling"])
+@pytest.mark.parametrize(
+ ("kernel_computation", "cloud_computation", "message"),
+ [
+ (True, False, "kernel computation"),
+ (False, True, "cloud computation"),
+ ],
+)
+def test_component_static_app_exports_reject_live_computation(
+ render_type, kernel_computation, cloud_computation, message
+):
+ component = Component(
+ walker=SimpleNamespace(kernel_computation=kernel_computation, cloud_computation=cloud_computation),
+ render_type=render_type,
+ field_map={},
+ single_chart_spec={},
+ )
+
+ with pytest.raises(ValueError, match=message):
+ component.to_html()
@pytest.mark.parametrize(
@@ -300,11 +1823,34 @@ def fake_webserver_walk(*args, **kwargs):
result = adapter.walk(
pd.DataFrame([{"city": "London", "value": 1}]),
gid="entry",
- kernel_computation=True,
+ spec_path="adapter_spec.json",
+ computation="kernel",
)
assert result == f"{expected_backend}-walker"
assert [call[0] for call in calls] == [expected_backend]
+ assert calls[0][2]["spec_path"] == "adapter_spec.json"
+ assert calls[0][2]["computation"] == "kernel"
if expected_backend == "webserver":
assert calls[0][2]["auto_open"] is True
assert calls[0][2]["auto_shutdown"] is True
+
+
+def test_public_walk_forwards_legacy_kernel_flag_to_webserver(monkeypatch):
+ calls = []
+
+ def fake_webserver_walk(*args, **kwargs):
+ calls.append(("webserver", args, kwargs))
+ return "webserver-walker"
+
+ monkeypatch.setattr(adapter, "get_current_env", lambda: "script")
+ monkeypatch.setattr(adapter.webserver, "walk", fake_webserver_walk)
+
+ result = adapter.walk(
+ pd.DataFrame([{"city": "London", "value": 1}]),
+ gid="entry",
+ use_kernel_calc=True,
+ )
+
+ assert result == "webserver-walker"
+ assert calls[0][2]["use_kernel_calc"] is True
diff --git a/tests/test_pyproject_metadata.py b/tests/test_pyproject_metadata.py
new file mode 100644
index 00000000..9c4814ca
--- /dev/null
+++ b/tests/test_pyproject_metadata.py
@@ -0,0 +1,97 @@
+from pathlib import Path
+
+
+def _extract_extra(pyproject_text: str, extra_name: str) -> list[str]:
+ start = pyproject_text.index(f"{extra_name} = [")
+ end = pyproject_text.index("]", start)
+ return [line.strip().strip('",') for line in pyproject_text[start:end].splitlines()[1:] if line.strip()]
+
+
+def _extract_project_dependencies(pyproject_text: str) -> list[str]:
+ start = pyproject_text.index("dependencies = [")
+ end = pyproject_text.index("]", start)
+ return [line.strip().strip('",') for line in pyproject_text[start:end].splitlines()[1:] if line.strip()]
+
+
+def _extract_hatch_build_include(pyproject_text: str) -> list[str]:
+ section_start = pyproject_text.index("[tool.hatch.build]")
+ include_start = pyproject_text.index("include = [", section_start)
+ include_end = pyproject_text.index("]", include_start)
+ return [
+ line.strip().strip('",') for line in pyproject_text[include_start:include_end].splitlines()[1:] if line.strip()
+ ]
+
+
+def _extract_hatch_sdist_include(pyproject_text: str) -> list[str]:
+ section_start = pyproject_text.index("[tool.hatch.build.targets.sdist]")
+ include_start = pyproject_text.index("include = [", section_start)
+ include_end = pyproject_text.index("]", include_start)
+ return [
+ line.strip().strip('",') for line in pyproject_text[include_start:include_end].splitlines()[1:] if line.strip()
+ ]
+
+
+def test_jupyter_notebook_extra_targets_modern_widgets_by_default():
+ pyproject = (Path(__file__).resolve().parents[1] / "pyproject.toml").read_text(encoding="utf-8")
+
+ assert _extract_extra(pyproject, "notebook") == [
+ "jupyter-client>7.4.9",
+ "jupyter-server>2.5.0",
+ "ipywidgets>=8.0.0",
+ ]
+ assert _extract_extra(pyproject, "labv4") == _extract_extra(pyproject, "notebook")
+ assert _extract_extra(pyproject, "notebook-legacy") == [
+ "jupyter-client<=7.4.9,>6.0.0",
+ "jupyter-server<=2.5.0",
+ "ipywidgets<8.0.0,>7.0.0",
+ ]
+
+
+def test_project_metadata_declares_supported_python_and_bounded_dependencies():
+ pyproject = (Path(__file__).resolve().parents[1] / "pyproject.toml").read_text(encoding="utf-8")
+ dependencies = _extract_project_dependencies(pyproject)
+
+ assert 'requires-python = ">=3.10"' in pyproject
+ assert "requests>=2.31,<3" in dependencies
+ assert "sqlalchemy>=1.4,<3" in dependencies
+ assert "pydantic>=1.10,<3" in dependencies
+ assert "duckdb>=0.10.4,<2.0.0" in dependencies
+ assert "pyarrow>=10,<25" in dependencies
+
+
+def test_dev_extra_contains_ci_quality_tools():
+ pyproject = (Path(__file__).resolve().parents[1] / "pyproject.toml").read_text(encoding="utf-8")
+
+ assert {"pytest", "ruff", "nbmake"}.issubset(_extract_extra(pyproject, "dev"))
+
+
+def test_jupyter_builder_does_not_skip_existing_frontend_bundle():
+ pyproject = (Path(__file__).resolve().parents[1] / "pyproject.toml").read_text(encoding="utf-8")
+
+ assert "skip-if-exists" not in pyproject
+
+
+def test_package_declares_pep561_typed_marker():
+ repo_root = Path(__file__).resolve().parents[1]
+ pyproject = (repo_root / "pyproject.toml").read_text(encoding="utf-8")
+
+ assert (repo_root / "pygwalker/py.typed").is_file()
+ assert "pygwalker" in _extract_hatch_build_include(pyproject)
+
+
+def test_public_dataframe_type_alias_includes_pyarrow_table():
+ repo_root = Path(__file__).resolve().parents[1]
+ typing_source = (repo_root / "pygwalker/_typing.py").read_text(encoding="utf-8")
+
+ assert "import pyarrow as pa" in typing_source
+ assert "dataframe_types.append(pa.Table)" in typing_source
+
+
+def test_package_keeps_pygwalker_tools_metrics_namespace():
+ repo_root = Path(__file__).resolve().parents[1]
+ pyproject = (repo_root / "pyproject.toml").read_text(encoding="utf-8")
+
+ assert (repo_root / "pygwalker_tools/metrics").is_dir()
+ assert (repo_root / "tests/test_metrics_tools.py").is_file()
+ assert "pygwalker_tools" in _extract_hatch_build_include(pyproject)
+ assert "pygwalker_tools" in _extract_hatch_sdist_include(pyproject)
diff --git a/tests/test_readme_api_reference.py b/tests/test_readme_api_reference.py
new file mode 100644
index 00000000..958a9594
--- /dev/null
+++ b/tests/test_readme_api_reference.py
@@ -0,0 +1,169 @@
+import ast
+import inspect
+from pathlib import Path
+
+import pytest
+
+from pygwalker.api import jupyter
+
+
+TRANSLATED_README_NOTICE = (
+ "This translation is community-maintained and may lag behind the [English README](../README.md). "
+ "Treat the English README as the source of truth for API reference, installation, and development instructions."
+)
+
+
+REPO_ROOT = Path(__file__).resolve().parents[1]
+
+
+def _read_walk_api_table_rows() -> list[dict[str, str]]:
+ readme = REPO_ROOT / "README.md"
+ lines = readme.read_text(encoding="utf-8").splitlines()
+
+ heading_index = lines.index("### [pygwalker.walk](https://pygwalker-docs.vercel.app/api-reference/jupyter#walk)")
+ table_lines = []
+ for line in lines[heading_index + 1 :]:
+ if not line.startswith("|"):
+ if table_lines:
+ break
+ continue
+ table_lines.append(line)
+
+ rows = table_lines[2:]
+ parsed_rows = []
+ for row in rows:
+ cells = [cell.strip().strip("`") for cell in row.split("|")[1:-1]]
+ parsed_rows.append(
+ {
+ "parameter": cells[0],
+ "type": cells[1],
+ "default": cells[2],
+ "description": cells[3],
+ }
+ )
+ return parsed_rows
+
+
+def _read_walk_api_table_params() -> list[str]:
+ return [row["parameter"] for row in _read_walk_api_table_rows()]
+
+
+def _format_signature_default(parameter: inspect.Parameter) -> str:
+ if parameter.default is inspect.Parameter.empty:
+ return "-"
+ if parameter.default == "":
+ return '""'
+ return repr(parameter.default)
+
+
+def _read_source_docstring(relative_path: str, qualified_name: str) -> str:
+ source = (REPO_ROOT / relative_path).read_text(encoding="utf-8")
+ module = ast.parse(source)
+ parts = qualified_name.split(".")
+
+ if len(parts) == 1:
+ for node in module.body:
+ if isinstance(node, ast.FunctionDef) and node.name == parts[0]:
+ docstring = ast.get_docstring(node)
+ assert docstring is not None
+ return docstring
+ elif len(parts) == 2:
+ for node in module.body:
+ if isinstance(node, ast.ClassDef) and node.name == parts[0]:
+ for child in node.body:
+ if isinstance(child, ast.FunctionDef) and child.name == parts[1]:
+ docstring = ast.get_docstring(child)
+ assert docstring is not None
+ return docstring
+
+ raise AssertionError(f"Could not find docstring for {qualified_name} in {relative_path}")
+
+
+def test_readme_walk_api_table_matches_jupyter_walk_signature():
+ signature_params = []
+ for parameter in inspect.signature(jupyter.walk).parameters.values():
+ if parameter.kind is inspect.Parameter.VAR_KEYWORD:
+ signature_params.append(f"**{parameter.name}")
+ else:
+ signature_params.append(parameter.name)
+
+ assert _read_walk_api_table_params() == signature_params
+
+
+def test_readme_walk_api_table_defaults_match_jupyter_walk_signature():
+ table_defaults = {row["parameter"]: row["default"] for row in _read_walk_api_table_rows()}
+ signature_defaults = {}
+ for parameter in inspect.signature(jupyter.walk).parameters.values():
+ name = f"**{parameter.name}" if parameter.kind is inspect.Parameter.VAR_KEYWORD else parameter.name
+ signature_defaults[name] = _format_signature_default(parameter)
+
+ assert table_defaults == signature_defaults
+
+
+def test_readme_walk_api_table_documents_reusable_walker_input():
+ dataset_row = next(row for row in _read_walk_api_table_rows() if row["parameter"] == "dataset")
+
+ assert "Walker" in dataset_row["type"]
+ assert "pyarrow.Table" in dataset_row["type"]
+
+
+@pytest.mark.parametrize(
+ ("relative_path", "qualified_name", "expected_fragments"),
+ [
+ ("pygwalker/api/adapter.py", "walk", ("pyarrow.Table", "pygwalker.Walker")),
+ ("pygwalker/api/adapter.py", "render", ("pyarrow.Table",)),
+ ("pygwalker/api/adapter.py", "table", ("pyarrow.Table",)),
+ ("pygwalker/api/anywidget.py", "walk", ("pyarrow.Table", "pygwalker.Walker")),
+ ("pygwalker/api/jupyter.py", "walk", ("pyarrow.Table", "pygwalker.Walker")),
+ ("pygwalker/api/jupyter.py", "render", ("pyarrow.Table",)),
+ ("pygwalker/api/jupyter.py", "table", ("pyarrow.Table",)),
+ ("pygwalker/api/marimo.py", "walk", ("pyarrow.Table", "pygwalker.Walker")),
+ ("pygwalker/api/webserver.py", "walk", ("pyarrow.Table", "pygwalker.Walker")),
+ ("pygwalker/api/webserver.py", "render", ("pyarrow.Table",)),
+ ("pygwalker/api/webserver.py", "table", ("pyarrow.Table",)),
+ ("pygwalker/api/streamlit.py", "StreamlitRenderer.__init__", ("pyarrow.Table", "pygwalker.Walker")),
+ ("pygwalker/api/streamlit.py", "get_streamlit_html", ("pyarrow.Table", "pygwalker.Walker")),
+ ("pygwalker/api/gradio.py", "get_html_on_gradio", ("pyarrow.Table",)),
+ ("pygwalker/api/component.py", "component", ("pyarrow.Table",)),
+ ("pygwalker/api/kanaries_cloud.py", "create_cloud_dataset", ("pyarrow.Table",)),
+ ("pygwalker/api/kanaries_cloud.py", "create_cloud_walker", ("pyarrow.Table",)),
+ ("pygwalker/api/html.py", "to_html", ("pyarrow.Table", "pygwalker.Walker")),
+ ("pygwalker/api/html.py", "to_table_html", ("pyarrow.Table",)),
+ ("pygwalker/api/html.py", "to_render_html", ("pyarrow.Table",)),
+ ("pygwalker/api/html.py", "to_chart_html", ("pyarrow.Table",)),
+ ],
+)
+def test_public_api_docstrings_document_current_dataset_inputs(relative_path, qualified_name, expected_fragments):
+ docstring = _read_source_docstring(relative_path, qualified_name)
+
+ assert "pl.DataFrame | pd.DataFrame" not in docstring
+ for fragment in expected_fragments:
+ assert fragment in docstring
+
+
+def test_readme_walk_api_table_marks_legacy_jupyter_envs_deprecated():
+ env_row = next(row for row in _read_walk_api_table_rows() if row["parameter"] == "env")
+
+ assert "JupyterAnywidget" in env_row["description"]
+ assert "deprecated legacy transports" in env_row["description"]
+
+
+def test_translated_readmes_defer_api_reference_to_english_readme():
+ docs_dir = Path(__file__).resolve().parents[1] / "docs"
+
+ translated_readmes = sorted(docs_dir.glob("README.*.md"))
+ assert translated_readmes
+ for readme in translated_readmes:
+ assert TRANSLATED_README_NOTICE in readme.read_text(encoding="utf-8")
+
+
+def test_translated_readmes_do_not_repeat_stale_api_parameters():
+ docs_dir = Path(__file__).resolve().parents[1] / "docs"
+ stale_parameters = {"hide_data_source_config", "use_preview"}
+
+ translated_readmes = sorted(docs_dir.glob("README.*.md"))
+ assert translated_readmes
+ for readme in translated_readmes:
+ source = readme.read_text(encoding="utf-8")
+ for parameter in stale_parameters:
+ assert parameter not in source
diff --git a/tests/test_repo_hygiene.py b/tests/test_repo_hygiene.py
new file mode 100644
index 00000000..b67b8e69
--- /dev/null
+++ b/tests/test_repo_hygiene.py
@@ -0,0 +1,45 @@
+import subprocess
+from pathlib import Path
+
+import pytest
+
+
+FORBIDDEN_TRACKED_PATH_PARTS = (
+ "__pycache__/",
+ ".pytest_cache/",
+ ".ipynb_checkpoints/",
+ "pygwalker/templates/dist/",
+)
+FORBIDDEN_TRACKED_SUFFIXES = (
+ ".pyc",
+ ".pyo",
+ ".DS_Store",
+ ".whl",
+)
+
+
+def _tracked_files(repo_root: Path) -> list[str]:
+ result = subprocess.run(
+ ["git", "ls-files"],
+ cwd=repo_root,
+ check=False,
+ text=True,
+ capture_output=True,
+ )
+ if result.returncode != 0:
+ pytest.skip("repo hygiene check requires a git checkout")
+ return result.stdout.splitlines()
+
+
+def test_generated_cache_and_package_artifacts_are_not_tracked():
+ repo_root = Path(__file__).resolve().parents[1]
+
+ forbidden_files = [
+ path
+ for path in _tracked_files(repo_root)
+ if any(part in path for part in FORBIDDEN_TRACKED_PATH_PARTS)
+ or any(path.endswith(suffix) for suffix in FORBIDDEN_TRACKED_SUFFIXES)
+ or (path.startswith("dist/") and not path.endswith(".gitkeep"))
+ ]
+
+ assert forbidden_files == []
diff --git a/tests/test_spec_communication.py b/tests/test_spec_communication.py
new file mode 100644
index 00000000..32bbb014
--- /dev/null
+++ b/tests/test_spec_communication.py
@@ -0,0 +1,101 @@
+from types import SimpleNamespace
+
+from pygwalker import __version__
+from pygwalker.communications.protocol import SaveChartRequest, UpdateSpecRequest
+from pygwalker.services.spec_communication import SpecCommunicationService
+
+
+def _chart_payload(title="Updated chart"):
+ return {
+ "charts": [
+ {
+ "rowIndex": 0,
+ "colIndex": 0,
+ "data": "data:image/png;base64,abc",
+ "height": 100,
+ "width": 200,
+ "canvasHeight": 100,
+ "canvasWidth": 200,
+ }
+ ],
+ "singleChart": "data:image/png;base64,abc",
+ "nRows": 1,
+ "nCols": 1,
+ "title": title,
+ }
+
+
+def test_spec_communication_returns_latest_vis_spec():
+ walker = SimpleNamespace(vis_spec=[{"name": "Chart"}])
+
+ assert SpecCommunicationService(walker).get_latest_vis_spec({}) == {"visSpec": [{"name": "Chart"}]}
+
+
+def test_spec_communication_saves_chart_payload():
+ saved_payloads = []
+ walker = SimpleNamespace(spec_manager=SimpleNamespace(save_chart_payload=saved_payloads.append))
+
+ response = SpecCommunicationService(walker).save_chart(SaveChartRequest(**_chart_payload("Saved chart")))
+
+ assert response == {}
+ assert saved_payloads == [_chart_payload("Saved chart")]
+
+
+def test_spec_communication_updates_runtime_state_writes_back_and_refreshes_preview():
+ runtime_updates = []
+ write_backs = []
+ preview_renders = []
+ preview_tool = SimpleNamespace(async_render_gw_review=preview_renders.append)
+ walker = SimpleNamespace(
+ use_preview=True,
+ cloud_service=object(),
+ spec_manager=SimpleNamespace(
+ update_runtime_state=lambda **kwargs: runtime_updates.append(kwargs),
+ write_back=lambda cloud_service, version: write_backs.append((cloud_service, version)),
+ ),
+ _get_gw_preview_html=lambda: "preview
",
+ )
+ request = UpdateSpecRequest(
+ visSpec=[{"name": "Chart"}],
+ workflowList=[{"workflow": []}],
+ chartData=_chart_payload("Chart"),
+ )
+
+ response = SpecCommunicationService(walker, preview_tool).update_spec(request)
+
+ assert response == {}
+ assert runtime_updates == [
+ {
+ "vis_spec": [{"name": "Chart"}],
+ "workflow_list": [{"workflow": []}],
+ "chart_data": _chart_payload("Chart"),
+ "version": __version__,
+ }
+ ]
+ assert write_backs == [(walker.cloud_service, __version__)]
+ assert preview_renders == ["preview
"]
+
+
+def test_spec_communication_skips_preview_refresh_without_preview_tool():
+ runtime_updates = []
+ write_backs = []
+ walker = SimpleNamespace(
+ use_preview=True,
+ cloud_service=object(),
+ spec_manager=SimpleNamespace(
+ update_runtime_state=lambda **kwargs: runtime_updates.append(kwargs),
+ write_back=lambda cloud_service, version: write_backs.append((cloud_service, version)),
+ ),
+ _get_gw_preview_html=lambda: "preview
",
+ )
+ request = UpdateSpecRequest(
+ visSpec=[{"name": "Chart"}],
+ workflowList=[{"workflow": []}],
+ chartData=_chart_payload("Chart"),
+ )
+
+ response = SpecCommunicationService(walker).update_spec(request)
+
+ assert response == {}
+ assert len(runtime_updates) == 1
+ assert write_backs == [(walker.cloud_service, __version__)]
diff --git a/tests/test_spec_input.py b/tests/test_spec_input.py
new file mode 100644
index 00000000..6c3829f8
--- /dev/null
+++ b/tests/test_spec_input.py
@@ -0,0 +1,32 @@
+import os
+
+import pandas as pd
+import pytest
+
+from pygwalker.utils.spec import resolve_spec_input
+
+
+def test_resolve_spec_input_uses_legacy_spec_when_no_spec_path():
+ assert resolve_spec_input("{}", None) == "{}"
+
+
+def test_resolve_spec_input_accepts_explicit_spec_path(tmp_path):
+ path = tmp_path / "chart.json"
+
+ assert resolve_spec_input("", path) == os.fspath(path)
+
+
+def test_resolve_spec_input_accepts_pathlike_legacy_spec(tmp_path):
+ path = tmp_path / "chart.json"
+
+ assert resolve_spec_input(path, None) == os.fspath(path)
+
+
+def test_resolve_spec_input_rejects_spec_and_spec_path_together(tmp_path):
+ with pytest.raises(ValueError, match="Pass only one of `spec` or `spec_path`"):
+ resolve_spec_input("{}", tmp_path / "chart.json")
+
+
+def test_resolve_spec_input_rejects_array_like_spec_without_ambiguous_truth_value(tmp_path):
+ with pytest.raises(ValueError, match="Pass only one of `spec` or `spec_path`"):
+ resolve_spec_input(pd.Series(["{}"]), tmp_path / "chart.json")
diff --git a/tests/test_spec_json.py b/tests/test_spec_json.py
new file mode 100644
index 00000000..e3c164ed
--- /dev/null
+++ b/tests/test_spec_json.py
@@ -0,0 +1,37 @@
+import pytest
+
+from pygwalker.errors import PrivacyError
+from pygwalker.services.global_var import GlobalVarManager
+from pygwalker.services.spec import get_spec_json
+
+
+@pytest.mark.parametrize(
+ "spec",
+ [
+ "ksf://workspace/spec.json",
+ "https://example.test/spec.json",
+ "a" * 32,
+ ],
+)
+def test_get_spec_json_rejects_remote_sources_in_offline_mode(spec):
+ previous_privacy = GlobalVarManager.privacy
+ GlobalVarManager.privacy = "offline"
+
+ try:
+ with pytest.raises(PrivacyError, match="privacy policy"):
+ get_spec_json(spec)
+ finally:
+ GlobalVarManager.privacy = previous_privacy
+
+
+def test_get_spec_json_reports_invalid_json_file(tmp_path):
+ spec_path = tmp_path / "broken.json"
+ spec_path.write_text("{not-valid-json", encoding="utf-8")
+
+ with pytest.raises(ValueError, match="spec is not a valid json"):
+ get_spec_json(str(spec_path))
+
+
+def test_get_spec_json_rejects_ambiguous_long_file_name():
+ with pytest.raises(ValueError, match="Spec file name too long"):
+ get_spec_json("x" * 201)
diff --git a/tests/test_spec_manager.py b/tests/test_spec_manager.py
new file mode 100644
index 00000000..6ab41c63
--- /dev/null
+++ b/tests/test_spec_manager.py
@@ -0,0 +1,162 @@
+import json
+
+from pygwalker.services.spec_manager import SpecManager
+
+
+RAW_FIELDS = [
+ {
+ "fid": "city",
+ "name": "city",
+ "semanticType": "nominal",
+ "analyticType": "dimension",
+ },
+ {
+ "fid": "value",
+ "name": "value",
+ "semanticType": "quantitative",
+ "analyticType": "measure",
+ },
+]
+
+
+def _vis_spec(name="Chart 1"):
+ return [
+ {
+ "name": name,
+ "encodings": {
+ "dimensions": [],
+ "measures": [],
+ },
+ }
+ ]
+
+
+def _chart_payload(title="Chart 1"):
+ return {
+ "charts": [
+ {
+ "rowIndex": 0,
+ "colIndex": 0,
+ "data": "data:image/png;base64,abc",
+ "height": 100,
+ "width": 200,
+ "canvasHeight": 100,
+ "canvasWidth": 200,
+ }
+ ],
+ "singleChart": "data:image/png;base64,abc",
+ "nRows": 1,
+ "nCols": 1,
+ "title": title,
+ }
+
+
+def test_spec_manager_initializes_spec_state_and_fills_new_fields():
+ manager = SpecManager(
+ {
+ "config": _vis_spec(),
+ "chart_map": {},
+ "workflow_list": [{"workflow": []}],
+ "version": "0.5.0",
+ },
+ RAW_FIELDS,
+ )
+
+ assert manager.spec_type == "json_obj"
+ assert manager.spec_version == "0.5.0"
+ assert manager.workflow_list == [{"workflow": []}]
+ assert manager.chart_name_index_map == {"Chart 1": 0}
+ assert [field["fid"] for field in manager.vis_spec[0]["encodings"]["dimensions"]] == ["city"]
+ assert [field["fid"] for field in manager.vis_spec[0]["encodings"]["measures"]] == ["value"]
+
+
+def test_spec_manager_updates_runtime_state_and_saved_chart_map():
+ manager = SpecManager("", RAW_FIELDS)
+ workflow_list = [{"workflow": [{"type": "view"}]}]
+
+ manager.update_runtime_state(
+ vis_spec=_vis_spec("Updated chart"),
+ workflow_list=workflow_list,
+ chart_data=_chart_payload("Updated chart"),
+ version="0.6.0",
+ )
+
+ assert manager.spec_version == "0.6.0"
+ assert manager.workflow_list == workflow_list
+ assert manager.chart_list == ["Updated chart"]
+ assert manager.get_chart_by_name("Updated chart").title == "Updated chart"
+ assert manager.get_chart_index("Updated chart") == 0
+ assert manager.build_spec_obj("0.6.1") == {
+ "config": _vis_spec("Updated chart"),
+ "chart_map": {},
+ "version": "0.6.1",
+ "workflow_list": workflow_list,
+ }
+
+
+def test_spec_manager_writes_json_file_specs(tmp_path):
+ path = tmp_path / "spec.json"
+ path.write_text(json.dumps({"config": [], "chart_map": {}, "workflow_list": [], "version": "0.5.0"}))
+ manager = SpecManager(str(path), RAW_FIELDS)
+ manager.update_runtime_state(
+ vis_spec=_vis_spec("File chart"),
+ workflow_list=[{"workflow": [{"type": "view"}]}],
+ chart_data=_chart_payload("File chart"),
+ version="0.6.0",
+ )
+
+ manager.write_back(cloud_service=None, version="0.6.1")
+
+ assert json.loads(path.read_text()) == {
+ "config": _vis_spec("File chart"),
+ "chart_map": {},
+ "version": "0.6.1",
+ "workflow_list": [{"workflow": [{"type": "view"}]}],
+ }
+
+
+def test_spec_manager_writes_ksf_cloud_specs():
+ manager = SpecManager("", RAW_FIELDS)
+ manager.spec = "ksf://workspace/spec.json"
+ manager.spec_type = "json_ksf"
+ manager.update_runtime_state(
+ vis_spec=_vis_spec("Cloud chart"),
+ workflow_list=[{"workflow": [{"type": "view"}]}],
+ chart_data=_chart_payload("Cloud chart"),
+ version="0.6.0",
+ )
+ writes = []
+
+ class FakeCloudService:
+ def write_config_to_cloud(self, path, payload):
+ writes.append((path, json.loads(payload)))
+
+ manager.write_back(FakeCloudService(), version="0.6.1")
+
+ assert writes == [
+ (
+ "workspace/spec.json",
+ {
+ "config": _vis_spec("Cloud chart"),
+ "chart_map": {},
+ "version": "0.6.1",
+ "workflow_list": [{"workflow": [{"type": "view"}]}],
+ },
+ )
+ ]
+
+
+def test_spec_manager_loads_saved_chart_map():
+ chart_payload = _chart_payload("Saved chart")
+ manager = SpecManager(
+ {
+ "config": [],
+ "chart_map": {"Saved chart": chart_payload},
+ "workflow_list": [],
+ "version": "0.5.0",
+ },
+ RAW_FIELDS,
+ )
+
+ assert manager.chart_list == ["Saved chart"]
+ assert manager.get_chart_by_name("Saved chart").title == "Saved chart"
diff --git a/tests/test_walker_api.py b/tests/test_walker_api.py
new file mode 100644
index 00000000..d128882d
--- /dev/null
+++ b/tests/test_walker_api.py
@@ -0,0 +1,254 @@
+import sys
+from types import SimpleNamespace
+
+import pandas as pd
+import pytest
+
+import pygwalker
+from pygwalker.api import html
+from pygwalker.api import walker as walker_api
+
+
+class FakeCoreWalker:
+ instances = []
+
+ def __init__(self, **kwargs):
+ self.kwargs = kwargs
+ self.gid = kwargs["gid"] or "generated"
+ self.kernel_computation = kwargs["kernel_computation"]
+ self.cloud_computation = kwargs["cloud_computation"]
+ self.display_calls = []
+ FakeCoreWalker.instances.append(self)
+
+ def display_on_jupyter_use_widgets(self, iframe_width=None, iframe_height=None):
+ self.display_calls.append(("jupyter-widget", iframe_width, iframe_height))
+
+ def display_on_jupyter_use_anywidget(self):
+ self.display_calls.append(("jupyter-anywidget",))
+
+ def display_on_jupyter(self):
+ self.display_calls.append(("jupyter-inline",))
+
+ def display_on_convert_html(self):
+ self.display_calls.append(("jupyter-convert",))
+
+ def to_html(self, iframe_width=None, iframe_height=None):
+ return f"iframe:{iframe_width}:{iframe_height}"
+
+ def to_html_without_iframe(self):
+ return "html-without-iframe"
+
+
+@pytest.fixture(autouse=True)
+def reset_fake_core_walker():
+ FakeCoreWalker.instances = []
+ yield
+ FakeCoreWalker.instances = []
+
+
+def test_public_package_exports_walker():
+ assert pygwalker.Walker is walker_api.Walker
+
+
+def test_walker_getattr_does_not_recurse_before_core_is_assigned():
+ walker = object.__new__(walker_api.Walker)
+
+ with pytest.raises(AttributeError):
+ getattr(walker, "missing")
+
+
+def test_walker_builds_core_with_unified_options(monkeypatch, tmp_path):
+ monkeypatch.setattr(walker_api, "PygWalker", FakeCoreWalker)
+ spec_path = tmp_path / "gw_config.json"
+ spec_path.write_text('{"config":[],"chart_map":{},"workflow_list":[],"version":"0.5.0"}')
+
+ walker = walker_api.Walker(
+ pd.DataFrame([{"city": "London", "value": 1}]),
+ gid="api",
+ spec_path=str(spec_path),
+ spec_io_mode="rw",
+ computation="browser",
+ default_tab="data",
+ appearance="light",
+ )
+
+ core = FakeCoreWalker.instances[0]
+ assert walker.core is core
+ assert core.kwargs["gid"] == "api"
+ assert core.kwargs["spec"] == str(spec_path)
+ assert core.kwargs["kernel_computation"] is False
+ assert core.kwargs["cloud_computation"] is False
+ assert core.kwargs["use_save_tool"] is True
+ assert core.kwargs["default_tab"] == "data"
+ assert core.kwargs["appearance"] == "light"
+
+
+def test_walker_preserves_auto_kernel_detection_by_default(monkeypatch):
+ monkeypatch.setattr(walker_api, "PygWalker", FakeCoreWalker)
+
+ walker_api.Walker(pd.DataFrame([{"city": "London", "value": 1}]))
+
+ assert FakeCoreWalker.instances[0].kwargs["kernel_computation"] is None
+
+
+def test_walker_show_auto_uses_current_notebook_env(monkeypatch):
+ monkeypatch.setattr(walker_api, "PygWalker", FakeCoreWalker)
+ monkeypatch.setattr(walker_api, "get_current_env", lambda: "jupyter")
+
+ walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), computation="browser")
+ result = walker.show(iframe_width="640px", iframe_height="480px")
+
+ assert result is walker
+ assert FakeCoreWalker.instances[0].display_calls == [("jupyter-anywidget",)]
+
+
+def test_walker_show_accepts_legacy_jupyter_widget_alias(monkeypatch):
+ monkeypatch.setattr(walker_api, "PygWalker", FakeCoreWalker)
+
+ walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), computation="browser")
+ with pytest.warns(DeprecationWarning, match="legacy Jupyter transport"):
+ walker.show("JupyterWidget", iframe_width="640px", iframe_height="480px")
+
+ assert FakeCoreWalker.instances[0].display_calls == [("jupyter-widget", "640px", "480px")]
+
+
+def test_walker_show_accepts_legacy_jupyter_env_alias(monkeypatch):
+ monkeypatch.setattr(walker_api, "PygWalker", FakeCoreWalker)
+
+ walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), computation="browser")
+ with pytest.warns(DeprecationWarning, match="legacy Jupyter transport"):
+ walker.show("Jupyter")
+
+ assert FakeCoreWalker.instances[0].display_calls == [("jupyter-inline",)]
+
+
+def test_walker_show_warns_for_lowercase_legacy_jupyter_env(monkeypatch):
+ monkeypatch.setattr(walker_api, "PygWalker", FakeCoreWalker)
+
+ walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), computation="browser")
+ with pytest.warns(DeprecationWarning, match="legacy Jupyter transport"):
+ walker.show("jupyter-widget", iframe_width="640px", iframe_height="480px")
+
+ assert FakeCoreWalker.instances[0].display_calls == [("jupyter-widget", "640px", "480px")]
+
+
+def test_walker_show_accepts_jupyter_preview_alias(monkeypatch):
+ monkeypatch.setattr(walker_api, "PygWalker", FakeCoreWalker)
+ monkeypatch.setattr(
+ FakeCoreWalker,
+ "display_preview_on_jupyter",
+ lambda self: self.display_calls.append(("jupyter-preview",)),
+ raising=False,
+ )
+
+ walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), computation="browser")
+ walker.show("JupyterPreview")
+
+ assert FakeCoreWalker.instances[0].display_calls == [("jupyter-preview",)]
+
+
+def test_walker_show_auto_uses_webserver_outside_notebooks(monkeypatch):
+ starts = []
+ monkeypatch.setattr(walker_api, "PygWalker", FakeCoreWalker)
+ monkeypatch.setattr(walker_api, "get_current_env", lambda: "script")
+ monkeypatch.setattr(walker_api.webserver, "_start_server", lambda *args, **kwargs: starts.append((args, kwargs)))
+
+ walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), computation="browser")
+ walker.show(port=8765, auto_open=False, auto_shutdown=False)
+
+ assert starts == [((FakeCoreWalker.instances[0], 8765), {"auto_open": False, "auto_shutdown": False})]
+
+
+@pytest.mark.parametrize("kwargs", [{"computation": "kernel"}, {"computation": "cloud"}])
+def test_walker_static_html_rejects_live_computation(monkeypatch, kwargs):
+ monkeypatch.setattr(walker_api, "PygWalker", FakeCoreWalker)
+
+ walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), **kwargs)
+
+ with pytest.raises(ValueError, match="Static HTML export does not support"):
+ walker.to_html()
+
+
+def test_to_html_adapter_accepts_walker_object(monkeypatch):
+ monkeypatch.setattr(walker_api, "PygWalker", FakeCoreWalker)
+
+ walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), computation="browser")
+ rendered = html.to_html(walker, width="640px", height="480px")
+
+ assert rendered == "iframe:640px:480px"
+
+
+def test_to_html_adapter_rejects_rebuilding_walker_object(monkeypatch):
+ monkeypatch.setattr(walker_api, "PygWalker", FakeCoreWalker)
+
+ walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), computation="browser")
+
+ with pytest.raises(ValueError, match="cannot apply construction parameters: spec_path"):
+ html.to_html(walker, spec_path="other.json")
+
+
+def test_to_html_adapter_preserves_walker_live_computation_error(monkeypatch):
+ monkeypatch.setattr(walker_api, "PygWalker", FakeCoreWalker)
+
+ walker = walker_api.Walker(pd.DataFrame([{"city": "London"}]), computation="kernel")
+
+ with pytest.raises(ValueError, match="Static HTML export does not support kernel computation"):
+ html.to_html(walker)
+
+
+def test_walker_to_streamlit_reuses_constructor_options(monkeypatch):
+ renderer_calls = []
+ monkeypatch.setattr(walker_api, "PygWalker", FakeCoreWalker)
+
+ class FakeStreamlitRenderer:
+ def __init__(self, dataset, gid=None, **kwargs):
+ renderer_calls.append((dataset, gid, kwargs))
+
+ monkeypatch.setitem(
+ sys.modules,
+ "pygwalker.api.streamlit",
+ SimpleNamespace(StreamlitRenderer=FakeStreamlitRenderer),
+ )
+
+ dataset = pd.DataFrame([{"city": "London"}])
+ walker = walker_api.Walker(dataset, gid="streamlit", spec="{}", computation="browser")
+
+ renderer = walker.to_streamlit(key="value")
+
+ assert isinstance(renderer, FakeStreamlitRenderer)
+ assert renderer_calls[0][0] is dataset
+ assert renderer_calls[0][1] == "streamlit"
+ assert renderer_calls[0][2]["spec"] == "{}"
+ assert renderer_calls[0][2]["computation"] == "browser"
+ assert "cloud_computation" not in renderer_calls[0][2]
+ assert renderer_calls[0][2]["key"] == "value"
+
+
+def test_walker_to_streamlit_maps_legacy_cloud_computation_and_drops_kernel_flags(monkeypatch):
+ renderer_calls = []
+ monkeypatch.setattr(walker_api, "PygWalker", FakeCoreWalker)
+
+ class FakeStreamlitRenderer:
+ def __init__(self, dataset, gid=None, **kwargs):
+ renderer_calls.append((dataset, gid, kwargs))
+
+ monkeypatch.setitem(
+ sys.modules,
+ "pygwalker.api.streamlit",
+ SimpleNamespace(StreamlitRenderer=FakeStreamlitRenderer),
+ )
+
+ with pytest.warns(DeprecationWarning, match="deprecated") as warnings:
+ walker = walker_api.Walker(
+ pd.DataFrame([{"city": "London"}]),
+ cloud_computation=True,
+ kernel_computation=True,
+ use_kernel_calc=True,
+ )
+ walker.to_streamlit()
+
+ assert len(warnings) == 3
+ assert renderer_calls[0][2]["computation"] == "cloud"
+ assert renderer_calls[0][2]["kernel_computation"] is None
+ assert renderer_calls[0][2]["use_kernel_calc"] is None
+ assert "cloud_computation" not in renderer_calls[0][2]