From a5b97546ecee2734f1bb97f69a8f31288676d346 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Erik=20Bj=C3=A4reholt?= Date: Sat, 26 Sep 2026 23:39:09 +0200 Subject: [PATCH 1/7] feat(transform)!: match aw-server-rust's merge_events_by_keys, split_url_events and tag Decided in ActivityWatch/activitywatch#1466: aw-core adopts aw-server-rust's output shape and semantics for these transforms, and chunk_events_by_key is deprecated in both servers. - merge_events_by_keys keeps the first event's whole data (not only the merge keys) and drops events missing a key; [] merges into []. aw-webui relies on both: top titles are colored by $category and top URLs/browser titles by $domain, and its multidevice/Android queries expect events without a title to be dropped. - split_url_events: $domain is the host without port, userinfo or leading www., $path includes ;params, $params is the query string, $options and $identifier are gone, and non-URLs are left unchanged. - tag: names must be strings (a query with category-style list names is rejected), and $tags is sorted and deduplicated. - chunk_events_by_key: DeprecationWarning and a query log warning; behavior unchanged. BREAKING CHANGE: hand-written queries on aw-server that used $options, $identifier or the old $params, relied on merge_events_by_keys keeping events without the keys, or passed list names to tag, get different results (the same as aw-server-rust). --- aw_query/functions.py | 16 +++ aw_transform/chunk_events_by_key.py | 11 ++ aw_transform/classify.py | 13 ++- aw_transform/merge_events_by_keys.py | 54 ++++----- aw_transform/split_url_events.py | 81 ++++++++++---- tests/test_query2.py | 18 +++ tests/test_transforms.py | 161 +++++++++++++++++---------- 7 files changed, 245 insertions(+), 109 deletions(-) diff --git a/aw_query/functions.py b/aw_query/functions.py index 7c2ad397..446b81bd 100644 --- a/aw_query/functions.py +++ b/aw_query/functions.py @@ -1,3 +1,4 @@ +import logging from datetime import timedelta from functools import wraps from inspect import signature @@ -36,6 +37,8 @@ from .exceptions import QueryFunctionException +logger = logging.getLogger(__name__) + def _verify_bucket_exists(datastore, bucketname): if bucketname in datastore.buckets(): @@ -256,6 +259,10 @@ def q2_merge_subwatcher_fields( @q2_function(chunk_events_by_key) @q2_typecheck def q2_chunk_events_by_key(events: list, key: str) -> List[Event]: + logger.warning( + "chunk_events_by_key is deprecated and will be removed, " + "use merge_events_by_keys instead" + ) return chunk_events_by_key(events, key) @@ -357,6 +364,15 @@ def q2_categorize(events: list, classes: list): @q2_function(tag) @q2_typecheck def q2_tag(events: list, classes: list): + # Tag names are strings, like in aw-server-rust. Category-style list names + # belong to categorize (ActivityWatch/activitywatch#1466). + for entry in classes: + if not isinstance(entry, list) or len(entry) != 2: + raise QueryFunctionException("tag expects a list of [name, rule] pairs") + if not isinstance(entry[0], str): + raise QueryFunctionException( + f"tag name must be a string, got {type(entry[0]).__name__}: {entry[0]!r}" + ) try: classes = [(_cls, Rule(rule_dict)) for _cls, rule_dict in classes] except ValueError as exc: diff --git a/aw_transform/chunk_events_by_key.py b/aw_transform/chunk_events_by_key.py index a99993b3..78c7c767 100644 --- a/aw_transform/chunk_events_by_key.py +++ b/aw_transform/chunk_events_by_key.py @@ -1,4 +1,5 @@ import logging +import warnings from datetime import timedelta from typing import List @@ -13,7 +14,17 @@ def chunk_events_by_key( """ "Chunks" adjacent events together which have the same value for a key, and stores the original events in the :code:`subevents` key of the new event. + + .. deprecated:: + Use :func:`merge_events_by_keys` instead. Nothing first-party uses this, + aw-server-rust never supported ``subevents``, and it will be removed + (ActivityWatch/activitywatch#1466). """ + warnings.warn( + "chunk_events_by_key is deprecated, use merge_events_by_keys instead", + DeprecationWarning, + stacklevel=2, + ) chunked_events: List[Event] = [] for event in events: if key not in event.data: diff --git a/aw_transform/classify.py b/aw_transform/classify.py index 35a33b81..1da84a00 100644 --- a/aw_transform/classify.py +++ b/aw_transform/classify.py @@ -79,7 +79,16 @@ def _categorize_one(e: Event, classes: List[Tuple[Category, Rule]]) -> Event: return e +def _matching_tags(e: Event, classes: List[Tuple[Tag, Rule]]) -> List[Tag]: + # Sorted and deduplicated, like aw-server-rust (ActivityWatch/activitywatch#1466) + return sorted({_cls for _cls, rule in classes if rule.match(e)}) + + def tag(events: List[Event], classes: List[Tuple[Tag, Rule]]) -> List[Event]: + """ + Adds the names of all matching rules to ``$tags`` (sorted, without + duplicates). Unlike categories, an event can have several tags. + """ cache: Dict[str, List[Tag]] = {} for e in events: try: @@ -87,13 +96,13 @@ def tag(events: List[Event], classes: List[Tuple[Tag, Rule]]) -> List[Event]: except TypeError: key = str(id(e.data)) if key not in cache: - cache[key] = [_cls for _cls, rule in classes if rule.match(e)] + cache[key] = _matching_tags(e, classes) e.data["$tags"] = list(cache[key]) return events def _tag_one(e: Event, classes: List[Tuple[Tag, Rule]]) -> Event: - e.data["$tags"] = [_cls for _cls, rule in classes if rule.match(e)] + e.data["$tags"] = _matching_tags(e, classes) return e diff --git a/aw_transform/merge_events_by_keys.py b/aw_transform/merge_events_by_keys.py index 42bccf9b..c3650305 100644 --- a/aw_transform/merge_events_by_keys.py +++ b/aw_transform/merge_events_by_keys.py @@ -1,40 +1,42 @@ +import copy +import json import logging -from typing import List, Dict, Tuple +from typing import Dict, List from aw_core.models import Event logger = logging.getLogger(__name__) -def merge_events_by_keys(events, keys) -> List[Event]: +def merge_events_by_keys(events: List[Event], keys: List[str]) -> List[Event]: """ - Sums the duration of all events which share a value for a key and returns a new event for each value. + Merges all events that share the same values for all of ``keys``, whether + they are adjacent or not, summing their durations. - .. note: The result will be a list of events without timestamp since they are merged. + Each merged event keeps the timestamp and the whole ``data`` of the first + event in its group (not only the merge keys), so fields that are the same + across the group, like ``$category`` for ``["app", "title"]``, stay + available. Events missing any of the keys are dropped, and an empty key list + returns no events. This matches aw-server-rust (ActivityWatch/activitywatch#1466). """ - # Call recursively until all keys are consumed - if len(keys) < 1: - return events - merged_events: Dict[Tuple, Event] = {} + if not keys: + return [] + merged_events: Dict[str, Event] = {} for event in events: - composite_key: Tuple = () - for key in keys: - if key in event.data: - val = event["data"][key] - # Needed for when the value is a list, such as for categories - if isinstance(val, list): - val = tuple(val) - composite_key = composite_key + (val,) - if composite_key not in merged_events: + try: + values = [event.data[key] for key in keys] + except KeyError: + continue + # Group by the JSON values, like aw-server-rust (so 1 and 1.0 differ, + # and list values such as categories work). + composite_key = json.dumps(values, sort_keys=True, default=str) + merged = merged_events.get(composite_key) + if merged is None: merged_events[composite_key] = Event( - timestamp=event.timestamp, duration=event.duration, data={} + timestamp=event.timestamp, + duration=event.duration, + data=copy.deepcopy(event.data), ) - for key in keys: - if key in event.data: - merged_events[composite_key].data[key] = event.data[key] else: - merged_events[composite_key].duration += event.duration - result = [] - for key in merged_events: - result.append(Event(**merged_events[key])) - return result + merged.duration += event.duration + return list(merged_events.values()) diff --git a/aw_transform/split_url_events.py b/aw_transform/split_url_events.py index 8a4466f6..8d4df6f0 100644 --- a/aw_transform/split_url_events.py +++ b/aw_transform/split_url_events.py @@ -1,32 +1,69 @@ import logging -from typing import List - -from urllib.parse import urlparse +from typing import List, Optional +from urllib.parse import SplitResult, urlsplit from aw_core.models import Event logger = logging.getLogger(__name__) +# WHATWG "special" schemes: they always have a host (except file) and a path +# that is at least "/", and their host is case-insensitive. +_SPECIAL_SCHEMES = {"http", "https", "ws", "wss", "ftp", "file"} + + +def _host(parts: SplitResult) -> str: + """The host of a URL without userinfo and port (IPv6 keeps its brackets).""" + host = parts.netloc.rpartition("@")[2] + if host.startswith("["): + return host[: host.find("]") + 1] if "]" in host else host + host = host.partition(":")[0] + return host.lower() if parts.scheme in _SPECIAL_SCHEMES else host + + +def _split_url(url: str) -> Optional[dict]: + try: + parts = urlsplit(url.strip()) + except ValueError: + return None + if not parts.scheme: + return None # a relative URL or plain text, not something we can split + host = _host(parts) + special = parts.scheme in _SPECIAL_SCHEMES + if special and parts.scheme != "file" and not host: + return None # e.g. "http://", which aw-server-rust rejects as well + domain = host + while domain.startswith("www."): + domain = domain[4:] + path = parts.path + if special and not path: + path = "/" + return { + "$protocol": parts.scheme, + # For URLs without a host (e.g. file://, about:), fall back to the + # scheme so they don't all cluster as an empty string. + "$domain": domain or parts.scheme, + "$path": path, + "$params": parts.query, + } + def split_url_events(events: List[Event]) -> List[Event]: + """ + Adds ``$protocol``, ``$domain``, ``$path`` and ``$params`` to events with a + ``url``, the same way as aw-server-rust (ActivityWatch/activitywatch#1466): + + - ``$domain`` is the host without port, userinfo or leading ``www.`` (the + scheme when there's no host, like ``about:blank`` or ``file:///x``) + - ``$path`` is the path, including any ``;params`` segment + - ``$params`` is the query string (``b=c`` in ``/a?b=c``) + + Events whose ``url`` isn't a string or an absolute URL are left unchanged. + """ for event in events: - if "url" in event.data: - url = event.data["url"] - parsed_url = urlparse(url) - event.data["$protocol"] = parsed_url.scheme - netloc = parsed_url.netloc - if netloc: - domain = netloc[4:] if netloc[:4] == "www." else netloc - elif parsed_url.scheme: - # For URLs without a domain (e.g. file://, about:), - # use the scheme as domain so they don't all cluster as empty. - domain = parsed_url.scheme - else: - domain = "" - event.data["$domain"] = domain - event.data["$path"] = parsed_url.path - event.data["$params"] = parsed_url.params - event.data["$options"] = parsed_url.query - event.data["$identifier"] = parsed_url.fragment - # TODO: Parse user, port etc aswell + url = event.data.get("url") + if not isinstance(url, str): + continue + fields = _split_url(url) + if fields is not None: + event.data.update(fields) return events diff --git a/tests/test_query2.py b/tests/test_query2.py index 70b89643..d1adbbac 100644 --- a/tests/test_query2.py +++ b/tests/test_query2.py @@ -351,6 +351,24 @@ def test_query2_categorize_invalid_priority(): query(qname, example_query, starttime, endtime, ds) +def test_query2_tag_requires_string_names(): + """Category-style list names are rejected, like in aw-server-rust (#1466)""" + ds = mock_ds + starttime = iso8601.parse_date("1970-01-01") + endtime = iso8601.parse_date("1970-01-02") + example_query = """ + events = []; + RETURN = tag(events, [[["Work"], {"type": "regex", "regex": "Code"}]]); + """ + with pytest.raises(QueryFunctionException, match="string"): + query("asd", example_query, starttime, endtime, ds) + ok_query = """ + events = []; + RETURN = tag(events, [["Work", {"type": "regex", "regex": "Code"}]]); + """ + assert query("asd", ok_query, starttime, endtime, ds) == [] + + @pytest.mark.parametrize("datastore", param_datastore_objects()) def test_query2_transforms_dont_modify_other_variables(datastore): """Query variables have value semantics, like in aw-server-rust""" diff --git a/tests/test_transforms.py b/tests/test_transforms.py index b0bc6211..99cf2780 100644 --- a/tests/test_transforms.py +++ b/tests/test_transforms.py @@ -296,11 +296,11 @@ def test_merge_events_by_keys_1(): events = events + [e1] * 10 events = events + [e2] * 5 - # Check that an empty key list has no effect - assert merge_events_by_keys(events, []) == events + # An empty key list merges nothing into nothing, like aw-server-rust + assert merge_events_by_keys(events, []) == [] - # Check that trying to merge on unavailable key has no effect - assert len(merge_events_by_keys(events, ["unknown"])) == 1 + # Events missing a key are dropped + assert merge_events_by_keys(events, ["unknown"]) == [] result = merge_events_by_keys(events, ["label"]) result = sort_by_duration(result) @@ -336,6 +336,45 @@ def test_merge_events_by_keys_2(): assert result[2].duration == timedelta(seconds=8) +def test_merge_events_by_keys_keeps_first_payload(): + """Like aw-server-rust: the first event's data is kept, missing keys are dropped""" + now = datetime.now(timezone.utc) + events = [ + Event( + timestamp=now, + duration=10, + data={"app": "x", "title": "t1", "$category": ["Work"]}, + ), + Event( + timestamp=now + timedelta(seconds=10), + duration=5, + data={"app": "x", "title": "t2", "$category": ["Work"]}, + ), + Event( + timestamp=now + timedelta(seconds=20), duration=7, data={"title": "no app"} + ), + ] + result = merge_events_by_keys(events, ["app"]) + assert len(result) == 1 + assert result[0].data == {"app": "x", "title": "t1", "$category": ["Work"]} + assert result[0].timestamp == events[0].timestamp + assert result[0].duration == timedelta(seconds=15) + assert result[0].id is None + + # The merged event doesn't share mutable data with the input + result[0].data["$category"].append("MUTATED") + assert events[0].data["$category"] == ["Work"] + + # List values (like categories) can be merge keys, and 1 and 1.0 differ + by_cat = merge_events_by_keys(events, ["$category"]) + assert [e.duration for e in by_cat] == [timedelta(seconds=15)] + nums = [ + Event(timestamp=now, duration=1, data={"n": 1}), + Event(timestamp=now, duration=1, data={"n": 1.0}), + ] + assert len(merge_events_by_keys(nums, ["n"])) == 2 + + def test_chunk_events_by_key(): now = datetime.now(timezone.utc) events = [] @@ -346,7 +385,8 @@ def test_chunk_events_by_key(): e2 = Event(data=e2_data, timestamp=now, duration=timedelta(seconds=1)) e3 = Event(data=e3_data, timestamp=now, duration=timedelta(seconds=1)) events = [e1, e2, e3] - result = chunk_events_by_key(events, "label1") + with pytest.warns(DeprecationWarning): + result = chunk_events_by_key(events, "label1") print(len(result)) pprint(result) assert len(result) == 2 @@ -365,59 +405,52 @@ def test_chunk_events_by_key(): assert result[1].data["subevents"][0] == e3 +def _split(url): + e = Event(data={"url": url}, timestamp=datetime.now(timezone.utc), duration=1) + return split_url_events([e])[0].data + + def test_url_parse_event(): - now = datetime.now(timezone.utc) - e = Event( - data={"url": "http://asd.com/test/?a=1"}, - timestamp=now, - duration=timedelta(seconds=1), - ) - result = split_url_events([e]) - print(result) - assert result[0].data["$protocol"] == "http" - assert result[0].data["$domain"] == "asd.com" - assert result[0].data["$path"] == "/test/" - assert result[0].data["$params"] == "" - assert result[0].data["$options"] == "a=1" - assert result[0].data["$identifier"] == "" - - e2 = Event( - data={"url": "https://www.asd.asd.com/test/test2/meh;meh2?asd=2&asdf=3#id"}, - timestamp=now, - duration=timedelta(seconds=1), - ) - result = split_url_events([e2]) - print(result) - assert result[0].data["$protocol"] == "https" - assert result[0].data["$domain"] == "asd.asd.com" - assert result[0].data["$path"] == "/test/test2/meh" - assert result[0].data["$params"] == "meh2" - assert result[0].data["$options"] == "asd=2&asdf=3" - assert result[0].data["$identifier"] == "id" - - e3 = Event( - data={"url": "file:///home/johan/myfile.txt"}, - timestamp=now, - duration=timedelta(seconds=1), - ) - result = split_url_events([e3]) - print(result) - assert result[0].data["$protocol"] == "file" - assert result[0].data["$domain"] == "file" - assert result[0].data["$path"] == "/home/johan/myfile.txt" - assert result[0].data["$params"] == "" - assert result[0].data["$options"] == "" - assert result[0].data["$identifier"] == "" - - # Test about: URLs - e4 = Event( - data={"url": "about:blank"}, - timestamp=now, - duration=timedelta(seconds=1), + """Same fields and values as aw-server-rust (ActivityWatch/activitywatch#1466)""" + assert _split("http://asd.com/test/?a=1") == { + "url": "http://asd.com/test/?a=1", + "$protocol": "http", + "$domain": "asd.com", + "$path": "/test/", + "$params": "a=1", + } + + data = _split("https://www.asd.asd.com/test/test2/meh;meh2?asd=2&asdf=3#id") + assert data["$domain"] == "asd.asd.com" + assert data["$path"] == "/test/test2/meh;meh2" + assert data["$params"] == "asd=2&asdf=3" + assert "$options" not in data and "$identifier" not in data + + # The port and userinfo aren't part of the domain + data = _split("https://user:pw@x.org:8080/a;p?b=c#d") + assert (data["$domain"], data["$path"], data["$params"]) == ("x.org", "/a;p", "b=c") + assert _split("http://[::1]:5600/api")["$domain"] == "[::1]" + assert _split("HTTPS://WWW.Example.COM")["$domain"] == "example.com" + + # Special schemes always have a path + assert _split("https://x.org")["$path"] == "/" + + # No host: the scheme is the domain + data = _split("file:///home/johan/myfile.txt") + assert (data["$protocol"], data["$domain"]) == ("file", "file") + assert data["$path"] == "/home/johan/myfile.txt" + data = _split("about:blank") + assert (data["$protocol"], data["$domain"], data["$path"]) == ( + "about", + "about", + "blank", ) - result = split_url_events([e4]) - assert result[0].data["$protocol"] == "about" - assert result[0].data["$domain"] == "about" + + # Not an absolute URL, or not a string: left unchanged + for url in ["not a url", "/relative/path", "http://"]: + assert _split(url) == {"url": url} + e = Event(data={"url": 5}, timestamp=datetime.now(timezone.utc), duration=1) + assert split_url_events([e])[0].data == {"url": 5} def test_union(): @@ -632,8 +665,17 @@ def test_tags(): ] events = tag(events, classes) - assert len(events[0].data["$tags"]) == 2 - assert len(events[1].data["$tags"]) == 0 + # Two rules with the same tag give it once + assert events[0].data["$tags"] == ["Test"] + assert events[1].data["$tags"] == [] + + # Sorted, not in rule order (like aw-server-rust) + classes = [ + ("Work", Rule({"regex": "Terminal"})), + ("Comms", Rule({"regex": "Inbox"})), + ] + e = Event(timestamp=now, duration=0, data={"app": "Terminal", "title": "Inbox"}) + assert tag([e], classes)[0].data["$tags"] == ["Comms", "Work"] def test_union_no_overlap(): @@ -847,7 +889,8 @@ def test_merge_subwatcher_fields_multiple_subsegments_preserve_duration(): assert by_project["alpha"] == 2 * td15m assert by_project["beta"] == td15m - assert by_project[None] == td15m + # merge_events_by_keys drops events without the key, so check those directly + assert sum_durations([e for e in result if "project" not in e.data]) == td15m assert by_app["vim"] == td1h assert sum_durations(result) == td1h From bd534e3fc0ee804beeb7c2981aa3b8020d1ebf3c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Erik=20Bj=C3=A4reholt?= Date: Sun, 27 Sep 2026 01:54:02 +0200 Subject: [PATCH 2/7] fix(transform): parse special-scheme URLs like the URL Standard in split_url_events --- aw_transform/split_url_events.py | 209 ++++++++++++++++++++++++++----- 1 file changed, 181 insertions(+), 28 deletions(-) diff --git a/aw_transform/split_url_events.py b/aw_transform/split_url_events.py index 8d4df6f0..d7284042 100644 --- a/aw_transform/split_url_events.py +++ b/aw_transform/split_url_events.py @@ -1,52 +1,201 @@ +import ipaddress import logging +import re from typing import List, Optional -from urllib.parse import SplitResult, urlsplit +from urllib.parse import unquote_to_bytes, urlsplit from aw_core.models import Event logger = logging.getLogger(__name__) -# WHATWG "special" schemes: they always have a host (except file) and a path -# that is at least "/", and their host is case-insensitive. +# The fields match aw-server-rust, which parses URLs with the WHATWG URL +# Standard (the `url` crate). For "special" schemes the parsing below follows +# the standard's steps that matter for these fields; other schemes have opaque +# hosts and paths and are split as-is. _SPECIAL_SCHEMES = {"http", "https", "ws", "wss", "ftp", "file"} +_SCHEME_RE = re.compile(r"([A-Za-z][A-Za-z0-9+.\-]*):(.*)", re.DOTALL) +# Characters a domain may not contain (after percent-decoding) +_FORBIDDEN_HOST = set(" #%/:<>?@[\\]^|") | {chr(c) for c in range(0x20)} | {"\x7f"} +# Percent-encode sets from the URL Standard +_PATH_ENCODE = set(' "#<>?`{}') +_QUERY_ENCODE = set(" \"#<>'") # "'" only for special schemes, which is all we encode -def _host(parts: SplitResult) -> str: - """The host of a URL without userinfo and port (IPv6 keeps its brackets).""" - host = parts.netloc.rpartition("@")[2] - if host.startswith("["): - return host[: host.find("]") + 1] if "]" in host else host - host = host.partition(":")[0] - return host.lower() if parts.scheme in _SPECIAL_SCHEMES else host +def _percent_encode(text: str, encode_set: set) -> str: + out = [] + for char in text: + if ord(char) < 0x21 or ord(char) > 0x7E or char in encode_set: + out.extend(f"%{byte:02X}" for byte in char.encode("utf-8")) + else: + out.append(char) + return "".join(out) -def _split_url(url: str) -> Optional[dict]: +def _is_dot(segment: str, dots: str) -> bool: + return segment.replace("%2e", ".").replace("%2E", ".") == dots + + +def _normalize_path(path: str) -> str: + """Special-scheme path: "/" separators, dot segments resolved, encoded.""" + segments = path.replace("\\", "/").split("/")[1:] + out: List[str] = [] + for i, segment in enumerate(segments): + last = i == len(segments) - 1 + if _is_dot(segment, ".."): + if out: + out.pop() + if last: + out.append("") + elif _is_dot(segment, "."): + if last: + out.append("") + else: + out.append(segment) + return "/" + "/".join(_percent_encode(s, _PATH_ENCODE) for s in out) + + +def _ipv4_number(part: str) -> Optional[int]: + if part[:2].lower() == "0x": + digits, base = part[2:], 16 + elif len(part) > 1 and part.startswith("0"): + digits, base = part[1:], 8 + else: + digits, base = part, 10 + if digits == "": + return 0 try: - parts = urlsplit(url.strip()) + return int(digits, base) except ValueError: return None - if not parts.scheme: - return None # a relative URL or plain text, not something we can split - host = _host(parts) - special = parts.scheme in _SPECIAL_SCHEMES - if special and parts.scheme != "file" and not host: - return None # e.g. "http://", which aw-server-rust rejects as well - domain = host - while domain.startswith("www."): - domain = domain[4:] + + +def _parse_ipv4(host: str) -> Optional[str]: + """IPv4 in any form the URL Standard accepts (127.1, 0x7f.1, 2130706433), + serialized dotted-decimal; "" if host isn't IPv4, None if it's invalid.""" + parts = host.split(".") + if parts[-1] == "" and len(parts) > 1: + parts.pop() + last = parts[-1] + ends_in_number = last.isdigit() or ( + last[:2].lower() == "0x" and all(c in "0123456789abcdefABCDEF" for c in last[2:]) + ) + if not ends_in_number: + return "" + if len(parts) > 4 or "" in parts: + return None + numbers = [_ipv4_number(p) for p in parts] + if any(n is None for n in numbers) or any(n > 255 for n in numbers[:-1]): + return None + if numbers[-1] >= 256 ** (5 - len(numbers)): + return None + value = numbers[-1] + for i, n in enumerate(numbers[:-1]): + value += n * 256 ** (3 - i) + return ".".join(str((value >> shift) & 0xFF) for shift in (24, 16, 8, 0)) + + +def _parse_host(host: str) -> Optional[str]: + """A special-scheme host as serialized by the URL Standard, or None if invalid.""" + if host.startswith("["): + if not host.endswith("]"): + return None + try: + return f"[{ipaddress.IPv6Address(host[1:-1]).compressed}]" + except ValueError: + return None + try: + decoded = unquote_to_bytes(host).decode("utf-8") + except UnicodeDecodeError: + return None + if not decoded or any(c in _FORBIDDEN_HOST for c in decoded): + return None + if decoded.isascii(): + ascii_host = decoded.lower() + else: + try: + ascii_host = decoded.encode("idna").decode("ascii").lower() + except UnicodeError: + return None + ipv4 = _parse_ipv4(ascii_host) + if ipv4 is None: + return None + return ipv4 or ascii_host + + +def _split_special(scheme: str, rest: str) -> Optional[dict]: + # Any number of slashes or backslashes may precede the authority. + rest = rest.lstrip("/\\") + end = len(rest) + for sep in "/\\?#": + idx = rest.find(sep) + if idx != -1: + end = min(end, idx) + authority, remainder = rest[:end], rest[end:] + hostport = authority.rpartition("@")[2] + if hostport.startswith("["): + close = hostport.find("]") + host, port = hostport[: close + 1], hostport[close + 1 :] + if port and not port.startswith(":"): + return None + port = port[1:] + else: + host, _, port = hostport.partition(":") + if port and (not port.isdigit() or int(port) > 65535): + return None + parsed_host = _parse_host(host) + if parsed_host is None: + return None + path, _, query = remainder.partition("#")[0].partition("?") + return { + "$protocol": scheme, + "$domain": _strip_www(parsed_host), + "$path": _normalize_path(path or "/"), + "$params": _percent_encode(query, _QUERY_ENCODE), + } + + +def _strip_www(host: str) -> str: + while host.startswith("www."): + host = host[4:] + return host + + +def _split_other(scheme: str, url: str) -> dict: + parts = urlsplit(url) + host = parts.netloc.rpartition("@")[2] + if not host.startswith("["): + host = host.partition(":")[0] path = parts.path - if special and not path: - path = "/" + if scheme == "file": + path = _normalize_path(path.replace("\\", "/") or "/") + host = host.lower() return { - "$protocol": parts.scheme, - # For URLs without a host (e.g. file://, about:), fall back to the - # scheme so they don't all cluster as an empty string. - "$domain": domain or parts.scheme, + "$protocol": scheme, + # No host (e.g. about:blank, file:///x): the scheme is the domain, so + # these don't all cluster as an empty string. + "$domain": _strip_www(host) or scheme, "$path": path, "$params": parts.query, } +def _split_url(url: str) -> Optional[dict]: + # Like the URL Standard: strip leading/trailing C0 controls and spaces, + # and remove tabs and newlines anywhere. + url = url.strip("".join(chr(c) for c in range(0x21))) + url = re.sub(r"[\t\n\r]", "", url) + match = _SCHEME_RE.fullmatch(url) + if not match: + return None # a relative URL or plain text, not something we can split + scheme, rest = match.group(1).lower(), match.group(2) + if scheme in _SPECIAL_SCHEMES and scheme != "file": + return _split_special(scheme, rest) + try: + return _split_other(scheme, f"{scheme}:{rest}") + except ValueError: + return None + + def split_url_events(events: List[Event]) -> List[Event]: """ Adds ``$protocol``, ``$domain``, ``$path`` and ``$params`` to events with a @@ -57,7 +206,11 @@ def split_url_events(events: List[Event]) -> List[Event]: - ``$path`` is the path, including any ``;params`` segment - ``$params`` is the query string (``b=c`` in ``/a?b=c``) - Events whose ``url`` isn't a string or an absolute URL are left unchanged. + For http(s), ws(s) and ftp URLs, the host is normalized (lowercase, + punycode, compressed IPv6) and the path and query are percent-encoded with + dot segments resolved, as the URL Standard does. Events whose ``url`` isn't + a string or a valid absolute URL (no scheme, an invalid host or port) are + left unchanged. Other fields of the event are never removed. """ for event in events: url = event.data.get("url") From 415dbe1eb2f6ab07df7cfc62d090bf595f0b49d9 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Erik=20Bj=C3=A4reholt?= Date: Sun, 27 Sep 2026 01:56:01 +0200 Subject: [PATCH 3/7] fix(query): address review of the #1466 changes - split_url_events: parse http(s)/ws(s)/ftp URLs like the URL Standard (aw-server-rust's url crate): any slashes or backslashes before the authority (http:example.com), invalid ports and hosts leave the event unchanged, hosts normalized (punycode, IPv4 forms, compressed IPv6), dot segments resolved and path/query percent-encoded. 117/117 URLs give the same fields as aw-server-rust's split_url_event. - merge_events_by_keys: tag non-JSON values (only possible through the Python API) with their type, so they can't merge with a lookalike string. - tag/categorize: a malformed rule is a QueryFunctionException instead of an AttributeError crash, like aw-server-rust's query error. - chunk_events_by_key: log the deprecation once per process, and don't suggest merge_events_by_keys as a replacement (it isn't one). --- aw_query/functions.py | 61 +++++++++++++++++++--------- aw_transform/chunk_events_by_key.py | 20 +++++---- aw_transform/merge_events_by_keys.py | 13 +++++- aw_transform/split_url_events.py | 10 +++-- tests/test_query2.py | 40 ++++++++++++++++++ tests/test_transforms.py | 54 ++++++++++++++++++++++++ 6 files changed, 164 insertions(+), 34 deletions(-) diff --git a/aw_query/functions.py b/aw_query/functions.py index 446b81bd..45c95b6a 100644 --- a/aw_query/functions.py +++ b/aw_query/functions.py @@ -1,4 +1,5 @@ import logging +import warnings from datetime import timedelta from functools import wraps from inspect import signature @@ -13,6 +14,7 @@ import iso8601 from aw_core.models import Event from aw_datastore import Datastore +from aw_transform.chunk_events_by_key import CHUNK_DEPRECATION from aw_transform import ( Rule, categorize, @@ -39,6 +41,9 @@ logger = logging.getLogger(__name__) +# Logged once per process: queries run repeatedly (e.g. on every dashboard refresh) +_chunk_deprecation_logged = False + def _verify_bucket_exists(datastore, bucketname): if bucketname in datastore.buckets(): @@ -259,11 +264,14 @@ def q2_merge_subwatcher_fields( @q2_function(chunk_events_by_key) @q2_typecheck def q2_chunk_events_by_key(events: list, key: str) -> List[Event]: - logger.warning( - "chunk_events_by_key is deprecated and will be removed, " - "use merge_events_by_keys instead" - ) - return chunk_events_by_key(events, key) + global _chunk_deprecation_logged + if not _chunk_deprecation_logged: + _chunk_deprecation_logged = True + logger.warning(CHUNK_DEPRECATION) + with warnings.catch_warnings(): + # Logged once above; don't also emit the transform's DeprecationWarning + warnings.simplefilter("ignore", DeprecationWarning) + return chunk_events_by_key(events, key) """ @@ -351,14 +359,32 @@ def q2_nop(): """ +def _parse_rules(function: str, classes: list) -> list: + """[[name, rule_dict], ...] into (name, Rule) pairs, with a + QueryFunctionException (like aw-server-rust's query error) for malformed + entries instead of a crash.""" + rules = [] + for entry in classes: + if not isinstance(entry, list) or len(entry) != 2: + raise QueryFunctionException( + f"{function} expects a list of [name, rule] pairs, got {entry!r}" + ) + name, rule_dict = entry + if not isinstance(rule_dict, dict): + raise QueryFunctionException( + f"{function} rule must be a dict, got {type(rule_dict).__name__}: {rule_dict!r}" + ) + try: + rules.append((name, Rule(rule_dict))) + except ValueError as exc: + raise QueryFunctionException(str(exc)) from None + return rules + + @q2_function(categorize) @q2_typecheck def q2_categorize(events: list, classes: list): - try: - classes = [(_cls, Rule(rule_dict)) for _cls, rule_dict in classes] - except ValueError as exc: - raise QueryFunctionException(str(exc)) from None - return categorize(_copy_events(events), classes) + return categorize(_copy_events(events), _parse_rules("categorize", classes)) @q2_function(tag) @@ -366,15 +392,10 @@ def q2_categorize(events: list, classes: list): def q2_tag(events: list, classes: list): # Tag names are strings, like in aw-server-rust. Category-style list names # belong to categorize (ActivityWatch/activitywatch#1466). - for entry in classes: - if not isinstance(entry, list) or len(entry) != 2: - raise QueryFunctionException("tag expects a list of [name, rule] pairs") - if not isinstance(entry[0], str): + rules = _parse_rules("tag", classes) + for name, _ in rules: + if not isinstance(name, str): raise QueryFunctionException( - f"tag name must be a string, got {type(entry[0]).__name__}: {entry[0]!r}" + f"tag name must be a string, got {type(name).__name__}: {name!r}" ) - try: - classes = [(_cls, Rule(rule_dict)) for _cls, rule_dict in classes] - except ValueError as exc: - raise QueryFunctionException(str(exc)) from None - return tag(_copy_events(events), classes) + return tag(_copy_events(events), rules) diff --git a/aw_transform/chunk_events_by_key.py b/aw_transform/chunk_events_by_key.py index 78c7c767..c4062763 100644 --- a/aw_transform/chunk_events_by_key.py +++ b/aw_transform/chunk_events_by_key.py @@ -8,6 +8,13 @@ logger = logging.getLogger(__name__) +CHUNK_DEPRECATION = ( + "chunk_events_by_key is deprecated and will be removed. There is no drop-in " + "replacement: merge_events_by_keys merges all events with the same value " + "(also across gaps) and doesn't produce subevents." +) + + def chunk_events_by_key( events: List[Event], key: str, pulsetime: float = 5.0 ) -> List[Event]: @@ -16,15 +23,12 @@ def chunk_events_by_key( original events in the :code:`subevents` key of the new event. .. deprecated:: - Use :func:`merge_events_by_keys` instead. Nothing first-party uses this, - aw-server-rust never supported ``subevents``, and it will be removed - (ActivityWatch/activitywatch#1466). + Will be removed (ActivityWatch/activitywatch#1466). There is no + drop-in replacement: aw-server-rust never supported ``subevents``, and + :func:`merge_events_by_keys` merges all events with the same value, + also across gaps, instead of adjacent runs. """ - warnings.warn( - "chunk_events_by_key is deprecated, use merge_events_by_keys instead", - DeprecationWarning, - stacklevel=2, - ) + warnings.warn(CHUNK_DEPRECATION, DeprecationWarning, stacklevel=2) chunked_events: List[Event] = [] for event in events: if key not in event.data: diff --git a/aw_transform/merge_events_by_keys.py b/aw_transform/merge_events_by_keys.py index c3650305..575cea04 100644 --- a/aw_transform/merge_events_by_keys.py +++ b/aw_transform/merge_events_by_keys.py @@ -1,13 +1,22 @@ import copy import json import logging -from typing import Dict, List +from typing import Any, Dict, List from aw_core.models import Event logger = logging.getLogger(__name__) +def _non_json_key(value: Any) -> Dict[str, str]: + # Event data from the datastore is always JSON, like in aw-server-rust. + # Other values (only possible through the Python API) are tagged with their + # type, so e.g. a datetime can't merge with a string that looks the same. + return { + "$non-json": f"{type(value).__module__}.{type(value).__qualname__}:{value!r}" + } + + def merge_events_by_keys(events: List[Event], keys: List[str]) -> List[Event]: """ Merges all events that share the same values for all of ``keys``, whether @@ -29,7 +38,7 @@ def merge_events_by_keys(events: List[Event], keys: List[str]) -> List[Event]: continue # Group by the JSON values, like aw-server-rust (so 1 and 1.0 differ, # and list values such as categories work). - composite_key = json.dumps(values, sort_keys=True, default=str) + composite_key = json.dumps(values, sort_keys=True, default=_non_json_key) merged = merged_events.get(composite_key) if merged is None: merged_events[composite_key] = Event( diff --git a/aw_transform/split_url_events.py b/aw_transform/split_url_events.py index d7284042..e5ff41b9 100644 --- a/aw_transform/split_url_events.py +++ b/aw_transform/split_url_events.py @@ -22,7 +22,7 @@ def _percent_encode(text: str, encode_set: set) -> str: - out = [] + out: List[str] = [] for char in text: if ord(char) < 0x21 or ord(char) > 0x7E or char in encode_set: out.extend(f"%{byte:02X}" for byte in char.encode("utf-8")) @@ -77,14 +77,16 @@ def _parse_ipv4(host: str) -> Optional[str]: parts.pop() last = parts[-1] ends_in_number = last.isdigit() or ( - last[:2].lower() == "0x" and all(c in "0123456789abcdefABCDEF" for c in last[2:]) + last[:2].lower() == "0x" + and all(c in "0123456789abcdefABCDEF" for c in last[2:]) ) if not ends_in_number: return "" if len(parts) > 4 or "" in parts: return None - numbers = [_ipv4_number(p) for p in parts] - if any(n is None for n in numbers) or any(n > 255 for n in numbers[:-1]): + parsed = [_ipv4_number(p) for p in parts] + numbers = [n for n in parsed if n is not None] + if len(numbers) != len(parsed) or any(n > 255 for n in numbers[:-1]): return None if numbers[-1] >= 256 ** (5 - len(numbers)): return None diff --git a/tests/test_query2.py b/tests/test_query2.py index d1adbbac..98b6cbd9 100644 --- a/tests/test_query2.py +++ b/tests/test_query2.py @@ -410,6 +410,46 @@ def test_query2_transforms_dont_modify_other_variables(datastore): assert result[0]["data"]["app"] == "Slack" +@pytest.mark.parametrize( + "q", + [ + 'RETURN = tag([], [["Work", "not a rule"]]);', + 'RETURN = tag([], [["Work", 5]]);', + 'RETURN = tag([], [["Work"]]);', + 'RETURN = categorize([], [[["Work"], "not a rule"]]);', + ], +) +def test_query2_malformed_rules_are_query_errors(q): + """A malformed rule is a query error (like aw-server-rust), not a crash""" + with pytest.raises(QueryFunctionException): + query( + "asd", + q, + iso8601.parse_date("1970-01-01"), + iso8601.parse_date("1970-01-02"), + mock_ds, + ) + + +def test_query2_chunk_deprecation_logged_once(caplog): + """The deprecation warning is logged once per process, not on every query""" + import aw_query.functions as functions + + functions._chunk_deprecation_logged = False + q = 'RETURN = chunk_events_by_key([], "app");' + start, end = iso8601.parse_date("1970-01-01"), iso8601.parse_date("1970-01-02") + with caplog.at_level("WARNING", logger="aw_query.functions"): + for _ in range(3): + query("asd", q, start, end, mock_ds) + messages = [ + r.getMessage() + for r in caplog.records + if "chunk_events_by_key" in r.getMessage() + ] + assert len(messages) == 1 + assert "no drop-in replacement" in messages[0] + + @pytest.mark.parametrize("datastore", param_datastore_objects()) def test_query2_function_in_function(datastore): qname = "asd" diff --git a/tests/test_transforms.py b/tests/test_transforms.py index 99cf2780..40d96af7 100644 --- a/tests/test_transforms.py +++ b/tests/test_transforms.py @@ -375,6 +375,16 @@ def test_merge_events_by_keys_keeps_first_payload(): assert len(merge_events_by_keys(nums, ["n"])) == 2 +def test_merge_events_by_keys_non_json_values_dont_collide(): + """A non-JSON value doesn't merge with a string that looks the same""" + now = datetime.now(timezone.utc) + events = [ + Event(timestamp=now, duration=1, data={"k": datetime(2020, 1, 1)}), + Event(timestamp=now, duration=1, data={"k": "2020-01-01 00:00:00"}), + ] + assert len(merge_events_by_keys(events, ["k"])) == 2 + + def test_chunk_events_by_key(): now = datetime.now(timezone.utc) events = [] @@ -453,6 +463,50 @@ def test_url_parse_event(): assert split_url_events([e])[0].data == {"url": 5} +def test_url_parse_event_like_url_standard(): + """Special-scheme URLs are parsed like aw-server-rust's WHATWG parser""" + + def fields(url): + d = _split(url) + return (d["$domain"], d["$path"], d["$params"]) if "$domain" in d else None + + # Any number of slashes (or backslashes) before the authority + assert fields("http:example.com") == ("example.com", "/", "") + assert fields("http:/example.com/a") == ("example.com", "/a", "") + assert fields("http:///x") == ("x", "/", "") + assert fields("http:\\\\x.org\\a") == ("x.org", "/a", "") + # Invalid ports and hosts leave the event unchanged + for url in [ + "https://example.com:not-a-port/a", + "https://example.com:65536/a", + "http://a b.com/", + "https://x.org%2Fevil/", + "http://1.2.3.4.5/", + "http://1..2/", + ]: + assert fields(url) is None, url + assert fields("https://example.com:65535/a") == ("example.com", "/a", "") + # Host normalization: punycode, IPv4 forms, compressed IPv6 + assert fields("https://bücher.de/") == ("xn--bcher-kva.de", "/", "") + assert fields("http://0x7f.1/") == ("127.0.0.1", "/", "") + assert fields("https://[::ffff:1.2.3.4]/") == ("[::ffff:102:304]", "/", "") + # Path and query: dot segments resolved, percent-encoded + assert fields("https://x.org/a/./b/../c") == ("x.org", "/a/c", "") + assert fields("https://x.org/a/%2e%2E/b") == ("x.org", "/b", "") + assert fields("https://x.org/a b?q=a b&r=ä#f g") == ( + "x.org", + "/a%20b", + "q=a%20b&r=%C3%A4", + ) + # Other fields of the event are kept, like in aw-server-rust + e = Event( + data={"url": "https://x.org/", "$options": "old"}, + timestamp=datetime.now(timezone.utc), + duration=1, + ) + assert split_url_events([e])[0].data["$options"] == "old" + + def test_union(): now = datetime.now(timezone.utc) From 064f21ea89ad8fbe210b2112c4b5a96f408e827d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Erik=20Bj=C3=A4reholt?= Date: Sun, 27 Sep 2026 02:43:00 +0200 Subject: [PATCH 4/7] fix(transform): split_url_events: ASCII ports, UTS 46 hosts, file URLs like the URL Standard MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit From the second review round of #173: - A port must be ASCII digits: "²" crashed int() and aborted the whole transform, and full-width digits were accepted. Any ValueError now leaves the event unchanged instead of failing the query. - Non-ASCII hosts use UTS 46 non-transitional processing (as the URL Standard does) instead of Python's IDNA 2003 codec, so faß.de is xn--fa-hia.de, not fass.de; zero-width joiners follow CheckJoiners. - file: URLs follow the URL Standard: "localhost" is no host, Windows drive letters stay in the path (file:c:/foo, file:C|/x), and ".." never removes the drive. 155/155 URLs give the same fields as aw-server-rust's split_url_event. --- aw_transform/split_url_events.py | 111 ++++++++++++++++++++++++++----- tests/test_transforms.py | 17 +++++ 2 files changed, 112 insertions(+), 16 deletions(-) diff --git a/aw_transform/split_url_events.py b/aw_transform/split_url_events.py index e5ff41b9..7bce008b 100644 --- a/aw_transform/split_url_events.py +++ b/aw_transform/split_url_events.py @@ -1,6 +1,7 @@ import ipaddress import logging import re +import unicodedata from typing import List, Optional from urllib.parse import unquote_to_bytes, urlsplit @@ -35,14 +36,21 @@ def _is_dot(segment: str, dots: str) -> bool: return segment.replace("%2e", ".").replace("%2E", ".") == dots -def _normalize_path(path: str) -> str: - """Special-scheme path: "/" separators, dot segments resolved, encoded.""" +_DRIVE_RE = re.compile(r"[A-Za-z][:|]") + + +def _normalize_path(path: str, file: bool = False, pipe_drive: bool = True) -> str: + """Special-scheme path: "/" separators, dot segments resolved, encoded. + For file URLs a leading Windows drive letter is kept ("C|" becomes "C:") + and ".." never removes it.""" segments = path.replace("\\", "/").split("/")[1:] out: List[str] = [] for i, segment in enumerate(segments): last = i == len(segments) - 1 + if file and i == 0 and pipe_drive and _DRIVE_RE.fullmatch(segment): + segment = segment[0] + ":" if _is_dot(segment, ".."): - if out: + if out and not (file and len(out) == 1 and _DRIVE_RE.fullmatch(out[0])): out.pop() if last: out.append("") @@ -96,6 +104,43 @@ def _parse_ipv4(host: str) -> Optional[str]: return ".".join(str((value >> shift) & 0xFF) for shift in (24, 16, 8, 0)) +# UTS 46 non-transitional processing keeps these instead of mapping them +# (ß to "ss", ς to σ), which the URL Standard requires. +_DEVIATIONS = set("\u00df\u03c2\u200c\u200d") +_LABEL_SEPARATORS = str.maketrans({"\u3002": ".", "\uff0e": ".", "\uff61": "."}) + + +def _uts46_to_ascii(host: str) -> Optional[str]: + """Non-ASCII host to ASCII (punycode) like the URL Standard's domain to + ASCII: UTS 46 mapping (compatibility forms, case folding, deviations kept), + NFC, then punycode per label.""" + mapped = "".join( + c + if c in _DEVIATIONS + else unicodedata.normalize("NFKC", unicodedata.normalize("NFKC", c).casefold()) + for c in host.translate(_LABEL_SEPARATORS) + ) + labels = [] + for label in unicodedata.normalize("NFC", mapped).split("."): + # Zero-width joiners are only valid after a virama (UTS 46 CheckJoiners) + for i, c in enumerate(label): + if c in "\u200c\u200d" and ( + i == 0 or unicodedata.combining(label[i - 1]) != 9 + ): + return None + if label.isascii(): + labels.append(label) + continue + try: + labels.append("xn--" + label.encode("punycode").decode("ascii")) + except UnicodeError: + return None + result = ".".join(labels) + if any(c in _FORBIDDEN_HOST for c in result): + return None + return result + + def _parse_host(host: str) -> Optional[str]: """A special-scheme host as serialized by the URL Standard, or None if invalid.""" if host.startswith("["): @@ -111,13 +156,9 @@ def _parse_host(host: str) -> Optional[str]: return None if not decoded or any(c in _FORBIDDEN_HOST for c in decoded): return None - if decoded.isascii(): - ascii_host = decoded.lower() - else: - try: - ascii_host = decoded.encode("idna").decode("ascii").lower() - except UnicodeError: - return None + ascii_host = decoded.lower() if decoded.isascii() else _uts46_to_ascii(decoded) + if ascii_host is None: + return None ipv4 = _parse_ipv4(ascii_host) if ipv4 is None: return None @@ -142,7 +183,7 @@ def _split_special(scheme: str, rest: str) -> Optional[dict]: port = port[1:] else: host, _, port = hostport.partition(":") - if port and (not port.isdigit() or int(port) > 65535): + if port and (not port.isascii() or not port.isdigit() or int(port) > 65535): return None parsed_host = _parse_host(host) if parsed_host is None: @@ -162,15 +203,50 @@ def _strip_www(host: str) -> str: return host +def _split_file(rest: str) -> Optional[dict]: + """file: URLs like the URL Standard: "localhost" is no host, and a Windows + drive letter is part of the path, not the host.""" + rest = rest.split("#", 1)[0] + rest, _, query = rest.partition("?") + host = "" + # Like the url crate, "C|" becomes "C:" except right after an empty host + # ("file:///c|/x" keeps it). + pipe_drive = True + if len(rest) >= 2 and rest[0] in "/\\" and rest[1] in "/\\": + rest = rest[2:] + end = len(rest) + for sep in "/\\": + idx = rest.find(sep) + if idx != -1: + end = min(end, idx) + authority, path = rest[:end], rest[end:] + if not authority: + pipe_drive = False + if _DRIVE_RE.fullmatch(authority): + path = "/" + authority + path + elif authority: + parsed = _parse_host(authority) + if parsed is None: + return None + host = "" if parsed == "localhost" else parsed + elif rest[:1] in ("/", "\\"): + path = rest + else: + path = "/" + rest + return { + "$protocol": "file", + "$domain": _strip_www(host) or "file", + "$path": _normalize_path(path or "/", file=True, pipe_drive=pipe_drive), + "$params": _percent_encode(query, _QUERY_ENCODE), + } + + def _split_other(scheme: str, url: str) -> dict: parts = urlsplit(url) host = parts.netloc.rpartition("@")[2] if not host.startswith("["): host = host.partition(":")[0] path = parts.path - if scheme == "file": - path = _normalize_path(path.replace("\\", "/") or "/") - host = host.lower() return { "$protocol": scheme, # No host (e.g. about:blank, file:///x): the scheme is the domain, so @@ -190,11 +266,14 @@ def _split_url(url: str) -> Optional[dict]: if not match: return None # a relative URL or plain text, not something we can split scheme, rest = match.group(1).lower(), match.group(2) - if scheme in _SPECIAL_SCHEMES and scheme != "file": - return _split_special(scheme, rest) try: + if scheme == "file": + return _split_file(rest) + if scheme in _SPECIAL_SCHEMES: + return _split_special(scheme, rest) return _split_other(scheme, f"{scheme}:{rest}") except ValueError: + # Never let one odd URL abort the whole transform: leave it unchanged. return None diff --git a/tests/test_transforms.py b/tests/test_transforms.py index 40d96af7..f1f44020 100644 --- a/tests/test_transforms.py +++ b/tests/test_transforms.py @@ -498,6 +498,23 @@ def fields(url): "/a%20b", "q=a%20b&r=%C3%A4", ) + # Non-ASCII digits aren't a port (and don't crash the transform) + assert fields("https://example.com:\u00b2/a") is None + assert fields("https://example.org:\uff11\uff12/a") is None + # UTS 46 non-transitional: ß and ς are kept, compatibility forms mapped + assert fields("https://fa\u00df.de/") == ("xn--fa-hia.de", "/", "") + assert ( + fields("https://\uff25\uff38\uff21\uff2d\uff30\uff2c\uff25.com/")[0] + == "example.com" + ) + # file: URLs: localhost is no host, Windows drive letters stay in the path + assert fields("file://localhost/etc") == ("file", "/etc", "") + assert fields("file:c:/foo") == ("file", "/c:/foo", "") + assert fields("file:C|/x") == ("file", "/C:/x", "") + assert fields("file:///C:/../a") == ("file", "/C:/a", "") + assert fields("file://host.example/share/x") == ("host.example", "/share/x", "") + # "?" and "#" end the authority, as in the URL Standard + assert fields("http://example.com?x/y") == ("example.com", "/", "x/y") # Other fields of the event are kept, like in aw-server-rust e = Event( data={"url": "https://x.org/", "$options": "old"}, From 8e4c709c0ee545f16bfdd913eef68f828438bac8 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Erik=20Bj=C3=A4reholt?= Date: Sun, 27 Sep 2026 03:16:36 +0200 Subject: [PATCH 5/7] fix(transform): full UTS 46 host processing and URL-standard non-special URLs From the third review round of #173: - Hosts go through UTS 46 with the idna package (new dependency, pure Python, BSD, already in the ActivityWatch bundle via requests): mapping removes soft hyphens and rejects disallowed code points, and the URL Standard's label checks reject leading combining marks, bad joiners and invalid A-labels (xn--zz, or ones decoding to disallowed text). ASCII hosts are checked too. - Non-special schemes (foo://, mailto:, chrome://, data:) follow the URL Standard instead of urlsplit: ports validated (invalid ones leave the event unchanged), IPv6 hosts without the port, opaque hosts and paths percent-encoded, dot segments resolved, queries percent-encoded. 191/191 URLs give the same fields as aw-server-rust's split_url_event, on Python 3.8 and 3.11. --- aw_transform/split_url_events.py | 148 ++++++++++++++++++++++--------- poetry.lock | 17 +++- pyproject.toml | 1 + tests/test_transforms.py | 15 ++++ 4 files changed, 140 insertions(+), 41 deletions(-) diff --git a/aw_transform/split_url_events.py b/aw_transform/split_url_events.py index 7bce008b..6af6262a 100644 --- a/aw_transform/split_url_events.py +++ b/aw_transform/split_url_events.py @@ -3,7 +3,9 @@ import re import unicodedata from typing import List, Optional -from urllib.parse import unquote_to_bytes, urlsplit +from urllib.parse import unquote_to_bytes + +import idna from aw_core.models import Event @@ -39,11 +41,14 @@ def _is_dot(segment: str, dots: str) -> bool: _DRIVE_RE = re.compile(r"[A-Za-z][:|]") -def _normalize_path(path: str, file: bool = False, pipe_drive: bool = True) -> str: +def _normalize_path( + path: str, file: bool = False, pipe_drive: bool = True, special: bool = True +) -> str: """Special-scheme path: "/" separators, dot segments resolved, encoded. For file URLs a leading Windows drive letter is kept ("C|" becomes "C:") and ".." never removes it.""" - segments = path.replace("\\", "/").split("/")[1:] + # "\\" is a separator only in special URLs + segments = (path.replace("\\", "/") if special else path).split("/")[1:] out: List[str] = [] for i, segment in enumerate(segments): last = i == len(segments) - 1 @@ -104,37 +109,51 @@ def _parse_ipv4(host: str) -> Optional[str]: return ".".join(str((value >> shift) & 0xFF) for shift in (24, 16, 8, 0)) -# UTS 46 non-transitional processing keeps these instead of mapping them -# (ß to "ss", ς to σ), which the URL Standard requires. -_DEVIATIONS = set("\u00df\u03c2\u200c\u200d") -_LABEL_SEPARATORS = str.maketrans({"\u3002": ".", "\uff0e": ".", "\uff61": "."}) +def _valid_label(label: str) -> bool: + """The URL Standard's UTS 46 validity checks that apply to a mapped label.""" + if label.startswith("xn--"): + # An A-label must decode, to a label that is itself valid + try: + decoded = label[4:].encode("ascii").decode("punycode") + except UnicodeError: + return False + if not decoded or decoded.isascii(): + return False + try: + # The decoded label must already be in mapped form (and allowed) + remapped = idna.uts46_remap(decoded, std3_rules=False, transitional=False) + except idna.IDNAError: + return False + return remapped == decoded and _valid_label(decoded) + if not label or label.isascii(): + return True + if unicodedata.category(label[0]).startswith("M"): + return False # a label can't begin with a combining mark + for i, c in enumerate(label): + # Zero-width joiners are only valid after a virama (CheckJoiners) + if c in "\u200c\u200d" and unicodedata.combining(label[i - 1]) != 9: + return False + return unicodedata.is_normalized("NFC", label) def _uts46_to_ascii(host: str) -> Optional[str]: - """Non-ASCII host to ASCII (punycode) like the URL Standard's domain to - ASCII: UTS 46 mapping (compatibility forms, case folding, deviations kept), - NFC, then punycode per label.""" - mapped = "".join( - c - if c in _DEVIATIONS - else unicodedata.normalize("NFKC", unicodedata.normalize("NFKC", c).casefold()) - for c in host.translate(_LABEL_SEPARATORS) - ) + """The URL Standard's domain to ASCII: UTS 46 mapping (non-transitional, + so ß and ς are kept; soft hyphens removed; compatibility forms and case + mapped), validity checks, then punycode for non-ASCII labels. None if the + host isn't valid.""" + try: + mapped = idna.uts46_remap(host, std3_rules=False, transitional=False) + except idna.IDNAError: + return None labels = [] - for label in unicodedata.normalize("NFC", mapped).split("."): - # Zero-width joiners are only valid after a virama (UTS 46 CheckJoiners) - for i, c in enumerate(label): - if c in "\u200c\u200d" and ( - i == 0 or unicodedata.combining(label[i - 1]) != 9 - ): - return None - if label.isascii(): - labels.append(label) - continue - try: - labels.append("xn--" + label.encode("punycode").decode("ascii")) - except UnicodeError: + for label in mapped.split("."): + if not _valid_label(label): return None + labels.append( + label + if label.isascii() + else "xn--" + label.encode("punycode").decode("ascii") + ) result = ".".join(labels) if any(c in _FORBIDDEN_HOST for c in result): return None @@ -156,7 +175,8 @@ def _parse_host(host: str) -> Optional[str]: return None if not decoded or any(c in _FORBIDDEN_HOST for c in decoded): return None - ascii_host = decoded.lower() if decoded.isascii() else _uts46_to_ascii(decoded) + # ASCII hosts go through it too, so invalid "xn--" labels are rejected + ascii_host = _uts46_to_ascii(decoded) if ascii_host is None: return None ipv4 = _parse_ipv4(ascii_host) @@ -241,22 +261,70 @@ def _split_file(rest: str) -> Optional[dict]: } -def _split_other(scheme: str, url: str) -> dict: - parts = urlsplit(url) - host = parts.netloc.rpartition("@")[2] - if not host.startswith("["): - host = host.partition(":")[0] - path = parts.path +# Non-special schemes: opaque hosts and paths keep their case, and only the +# URL Standard's C0 control percent-encode set applies to opaque paths. +_OPAQUE_HOST_FORBIDDEN = _FORBIDDEN_HOST - {"%"} + + +def _split_other(scheme: str, rest: str) -> Optional[dict]: + rest = rest.split("#", 1)[0] + rest, _, query = rest.partition("?") + host = "" + if rest.startswith("//"): + rest = rest[2:] + slash = rest.find("/") + authority, path = (rest, "") if slash == -1 else (rest[:slash], rest[slash:]) + hostport = authority.rpartition("@")[2] + if hostport.startswith("["): + close = hostport.find("]") + if close == -1: + return None + host, port = hostport[: close + 1], hostport[close + 1 :] + if port and not port.startswith(":"): + return None + port = port[1:] + try: + host = f"[{ipaddress.IPv6Address(host[1:-1]).compressed}]" + except ValueError: + return None + else: + host, _, port = hostport.partition(":") + if any(c in _OPAQUE_HOST_FORBIDDEN for c in host): + return None + host = _percent_encode(host, set()) + if port and (not port.isascii() or not port.isdigit() or int(port) > 65535): + return None + path = _normalize_nonspecial_path(path) if path else "" + elif rest.startswith("/"): + path = _normalize_nonspecial_path(rest) + else: + # An opaque path (mailto:, about:, data:, javascript:) + path = _percent_encode_c0(rest) return { "$protocol": scheme, - # No host (e.g. about:blank, file:///x): the scheme is the domain, so - # these don't all cluster as an empty string. + # No host (e.g. about:blank): the scheme is the domain, so these don't + # all cluster as an empty string. "$domain": _strip_www(host) or scheme, "$path": path, - "$params": parts.query, + "$params": _percent_encode(query, _QUERY_ENCODE - {"'"}), } +def _percent_encode_c0(text: str) -> str: + out: List[str] = [] + for char in text: + if ord(char) < 0x20 or ord(char) > 0x7E: + out.extend(f"%{byte:02X}" for byte in char.encode("utf-8")) + else: + out.append(char) + return "".join(out) + + +def _normalize_nonspecial_path(path: str) -> str: + """A non-special path with segments: dot segments resolved, encoded.""" + return _normalize_path(path, special=False) + + def _split_url(url: str) -> Optional[dict]: # Like the URL Standard: strip leading/trailing C0 controls and spaces, # and remove tabs and newlines anywhere. @@ -271,7 +339,7 @@ def _split_url(url: str) -> Optional[dict]: return _split_file(rest) if scheme in _SPECIAL_SCHEMES: return _split_special(scheme, rest) - return _split_other(scheme, f"{scheme}:{rest}") + return _split_other(scheme, rest) except ValueError: # Never let one odd URL abort the whole transform: leave it unchanged. return None diff --git a/poetry.lock b/poetry.lock index 56bac177..775ca7ab 100644 --- a/poetry.lock +++ b/poetry.lock @@ -167,6 +167,21 @@ files = [ [package.extras] test = ["pytest (>=6)"] +[[package]] +name = "idna" +version = "3.15" +description = "Internationalized Domain Names in Applications (IDNA)" +optional = false +python-versions = ">=3.8" +groups = ["main"] +files = [ + {file = "idna-3.15-py3-none-any.whl", hash = "sha256:048adeaf8c2d788c40fee287673ccaa74c24ffd8dcf09ffa555a2fbb59f10ac8"}, + {file = "idna-3.15.tar.gz", hash = "sha256:ca962446ea538f7092a95e057da437618e886f4d349216d2b1e294abfdb65fdc"}, +] + +[package.extras] +all = ["mypy (>=1.11.2)", "pytest (>=8.3.2)", "ruff (>=0.6.2)"] + [[package]] name = "importlib-resources" version = "6.4.5" @@ -693,4 +708,4 @@ type = ["pytest-mypy"] [metadata] lock-version = "2.1" python-versions = "^3.8" -content-hash = "9ac93f2d6bac28473f4b5a4dabda037554159497caa7ae113ca25ea79f763cbe" +content-hash = "9881cc22bf799b6b8fdc764f15a3399b31a4ad70630cfd5e7ffe68543495b074" diff --git a/pyproject.toml b/pyproject.toml index 1293a5fc..c24ff00d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -31,6 +31,7 @@ tomlkit = "*" deprecation = "*" timeslot = "*" click = "*" +idna = ">=3.0" # UTS 46 host processing in split_url_events, as in the URL Standard [tool.poetry.group.dev.dependencies] mypy = "*" diff --git a/tests/test_transforms.py b/tests/test_transforms.py index f1f44020..43e74a15 100644 --- a/tests/test_transforms.py +++ b/tests/test_transforms.py @@ -507,6 +507,21 @@ def fields(url): fields("https://\uff25\uff38\uff21\uff2d\uff30\uff2c\uff25.com/")[0] == "example.com" ) + # UTS 46 validity: soft hyphens are removed, invalid labels are rejected + assert fields("https://foo\u00adbar.com/")[0] == "foobar.com" + assert fields("https://\u0301.com/") is None # leading combining mark + assert fields("https://xn--zz.com/") is None # invalid A-label + assert fields("https://xn--fa-hia.de/")[0] == "xn--fa-hia.de" + # Other schemes: ports validated, IPv6 without port, encoded like the URL Standard + assert fields("foo://host:not-a-port/a") is None + assert fields("foo://host:99999/a") is None + assert fields("foo://[::1]:80/a") == ("[::1]", "/a", "") + assert fields("foo://host/a b") == ("host", "/a%20b", "") + assert fields("mailto:user@example.com?subject=hello world") == ( + "mailto", + "user@example.com", + "subject=hello%20world", + ) # file: URLs: localhost is no host, Windows drive letters stay in the path assert fields("file://localhost/etc") == ("file", "/etc", "") assert fields("file:c:/foo") == ("file", "/c:/foo", "") From 0842e699dabaadf3025b709a215f7a4740922cd2 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Erik=20Bj=C3=A4reholt?= Date: Sun, 27 Sep 2026 03:26:06 +0200 Subject: [PATCH 6/7] fix(transform): CheckBidi, leading joiners and empty hosts in split_url_events From the fourth review round of #173: - Hosts with right-to-left text must satisfy RFC 5893 bidi rules in every label (UTS 46 CheckBidi, via idna.check_bidi), as in the URL Standard. - A joiner at the start of a label is invalid; the check read the label's last character for i == 0. - In non-special URLs, credentials or a port with an empty host (foo://:80/a, foo://@/a) are invalid and leave the event unchanged. 208/208 URLs give the same fields as aw-server-rust's split_url_event. --- aw_transform/split_url_events.py | 17 +++++++++++++++-- tests/test_transforms.py | 7 +++++++ 2 files changed, 22 insertions(+), 2 deletions(-) diff --git a/aw_transform/split_url_events.py b/aw_transform/split_url_events.py index 6af6262a..72a85c0a 100644 --- a/aw_transform/split_url_events.py +++ b/aw_transform/split_url_events.py @@ -131,7 +131,7 @@ def _valid_label(label: str) -> bool: return False # a label can't begin with a combining mark for i, c in enumerate(label): # Zero-width joiners are only valid after a virama (CheckJoiners) - if c in "\u200c\u200d" and unicodedata.combining(label[i - 1]) != 9: + if c in "\u200c\u200d" and (i == 0 or unicodedata.combining(label[i - 1]) != 9): return False return unicodedata.is_normalized("NFC", label) @@ -146,9 +146,18 @@ def _uts46_to_ascii(host: str) -> Optional[str]: except idna.IDNAError: return None labels = [] - for label in mapped.split("."): + split = mapped.split(".") + # CheckBidi (RFC 5893): in a domain with right-to-left text, every label + # must satisfy the bidi rules + bidi = any(unicodedata.bidirectional(c) in ("R", "AL", "AN") for c in mapped) + for label in split: if not _valid_label(label): return None + if bidi and label: + try: + idna.check_bidi(label, check_ltr=True) + except idna.IDNAError: + return None labels.append( label if label.isascii() @@ -275,6 +284,10 @@ def _split_other(scheme: str, rest: str) -> Optional[dict]: slash = rest.find("/") authority, path = (rest, "") if slash == -1 else (rest[:slash], rest[slash:]) hostport = authority.rpartition("@")[2] + # Credentials or a port need a host ("foo://:80/a", "foo://@/a") + if hostport in ("", ":") or hostport.startswith(":"): + if "@" in authority or ":" in hostport: + return None if hostport.startswith("["): close = hostport.find("]") if close == -1: diff --git a/tests/test_transforms.py b/tests/test_transforms.py index 43e74a15..dcd3166d 100644 --- a/tests/test_transforms.py +++ b/tests/test_transforms.py @@ -512,6 +512,13 @@ def fields(url): assert fields("https://\u0301.com/") is None # leading combining mark assert fields("https://xn--zz.com/") is None # invalid A-label assert fields("https://xn--fa-hia.de/")[0] == "xn--fa-hia.de" + # CheckBidi and CheckJoiners + assert fields("https://\u05d0\u05d1x.com/") is None # mixed direction label + assert fields("https://\u05d0\u05d1.com/")[0] == "xn--4dbc.com" + assert fields("https://\u200d\u094d.com/") is None # leading joiner + # Other schemes: credentials or a port need a host + assert fields("foo://:80/a") is None + assert fields("foo://@/a") is None # Other schemes: ports validated, IPv6 without port, encoded like the URL Standard assert fields("foo://host:not-a-port/a") is None assert fields("foo://host:99999/a") is None From 9558a55a8d0ce059a1eb56f3001ce47bdd3bc0e4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Erik=20Bj=C3=A4reholt?= Date: Sun, 27 Sep 2026 03:34:42 +0200 Subject: [PATCH 7/7] fix(transform): RFC 5892 joiner context and bidi on encoded labels From the fifth review round of #173: - CheckJoiners uses the full RFC 5892 CONTEXTJ rule (idna.valid_contextj), so a ZWNJ in a joining context (Persian hostnames) is accepted, as in aw-server-rust; the virama-only approximation rejected them. - CheckBidi looks at the Unicode form of every label, so an encoded mixed-direction label (xn--x-zhcd) is rejected too. 218/218 URLs give the same fields as aw-server-rust's split_url_event. --- aw_transform/split_url_events.py | 30 +++++++++++++++++++++++------- tests/test_transforms.py | 6 ++++++ 2 files changed, 29 insertions(+), 7 deletions(-) diff --git a/aw_transform/split_url_events.py b/aw_transform/split_url_events.py index 72a85c0a..1a2f3131 100644 --- a/aw_transform/split_url_events.py +++ b/aw_transform/split_url_events.py @@ -130,12 +130,20 @@ def _valid_label(label: str) -> bool: if unicodedata.category(label[0]).startswith("M"): return False # a label can't begin with a combining mark for i, c in enumerate(label): - # Zero-width joiners are only valid after a virama (CheckJoiners) - if c in "\u200c\u200d" and (i == 0 or unicodedata.combining(label[i - 1]) != 9): + # CheckJoiners: RFC 5892 CONTEXTJ (after a virama, or ZWNJ in a + # joining context, as in Persian) + if c in "\u200c\u200d" and not idna.valid_contextj(label, i): return False return unicodedata.is_normalized("NFC", label) +def _decode_label(label: str) -> str: + """The Unicode form of a (valid) label.""" + if label.startswith("xn--"): + return label[4:].encode("ascii").decode("punycode") + return label + + def _uts46_to_ascii(host: str) -> Optional[str]: """The URL Standard's domain to ASCII: UTS 46 mapping (non-transitional, so ß and ς are kept; soft hyphens removed; compatibility forms and case @@ -147,17 +155,25 @@ def _uts46_to_ascii(host: str) -> Optional[str]: return None labels = [] split = mapped.split(".") - # CheckBidi (RFC 5893): in a domain with right-to-left text, every label - # must satisfy the bidi rules - bidi = any(unicodedata.bidirectional(c) in ("R", "AL", "AN") for c in mapped) for label in split: if not _valid_label(label): return None - if bidi and label: + # CheckBidi (RFC 5893): in a domain with right-to-left text, every label + # must satisfy the bidi rules. Checked on the Unicode form, so encoded + # ("xn--") labels count too. + unicode_labels = [_decode_label(label) for label in split] + if any( + unicodedata.bidirectional(c) in ("R", "AL", "AN") + for label in unicode_labels + for c in label + ): + for label in unicode_labels: try: - idna.check_bidi(label, check_ltr=True) + if label: + idna.check_bidi(label, check_ltr=True) except idna.IDNAError: return None + for label in split: labels.append( label if label.isascii() diff --git a/tests/test_transforms.py b/tests/test_transforms.py index dcd3166d..537e6a12 100644 --- a/tests/test_transforms.py +++ b/tests/test_transforms.py @@ -516,6 +516,12 @@ def fields(url): assert fields("https://\u05d0\u05d1x.com/") is None # mixed direction label assert fields("https://\u05d0\u05d1.com/")[0] == "xn--4dbc.com" assert fields("https://\u200d\u094d.com/") is None # leading joiner + assert fields("https://xn--x-zhcd.com/") is None # encoded mixed-direction label + # ZWNJ in a joining context is valid (Persian) + persian = "https://\u0646\u0627\u0645\u0647\u200c\u0627\u06cc.com/" + assert fields(persian)[0] == "xn--mgba3gch31f060k.com" + # "^" is kept literally in paths, like aw-server-rust + assert fields("https://x.org/a^b") == ("x.org", "/a^b", "") # Other schemes: credentials or a port need a host assert fields("foo://:80/a") is None assert fields("foo://@/a") is None