diff --git a/CHANGELOG.md b/CHANGELOG.md index 057f8345..1984818c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,6 +5,23 @@ All notable changes to this project will be documented in this file. The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/), and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). +## [2.0.5] - 2026-10-05 + +### Improved + +- **`installQueries()` no longer holds the connection open while queries compile** — installation is submitted asynchronously, so the call returns as soon as the server accepts it instead of blocking for the whole compile and risking a gateway timeout. `wait=True` (the default for sync connections) still returns only once installation has finished; `wait=False` (the default for async connections) now returns the ID of the submitted request rather than a completion message. +- **`AI.query()` chat engine selection** — `query()` now accepts `mode` (`"agentic"` | `"classic"`), `rag_method` (agent style or retriever name), and `include_fields` to choose the GraphRAG chat engine and control what the response returns. Called with only a question it behaves exactly as before, deferring to the graph's configured default engine. + +### Fixed + +- **Declared minimum Python version corrected to 3.9** — installing on Python 3.8 previously succeeded and then failed at import. The package has in fact required 3.9 since it began using built-in generic type annotations; the conda recipe already required 3.9. +- **`AsyncTigerGraphConnection` use across event loops** — a connection is no longer tied to the first event loop it ran on. Reusing it after a loop has closed, which is what `asyncio.run()` leaves behind on every call, no longer fails with `RuntimeError: Event loop is closed`; and sharing one connection between threads that each run their own loop no longer lets one thread close the session another is still using. The HTTP session and connection pool are now kept per event loop, and those belonging to a closed loop are released rather than reported as leaked. +- **Untyped vertex IDs in `runInstalledQuery()` GET mode** — passing a vertex as `(id, "type")` no longer fails when the ID contains `&` or `#`; the ID is now escaped in the query string as it already was for every other vertex form. +- **Parameter values in async `runInterpretedQuery()`** — string and vertex values containing spaces, quotes, `&`, `=`, `%`, `#`, `+` or non-ASCII characters reached the query still percent-encoded (e.g. `"a b"` arrived as `"a%20b"`), and a vertex whose ID contained one of them could not be found. Values now arrive exactly as passed, matching the sync client. +- **Vertex IDs containing `%` in `runInstalledQuery()`** — a vertex whose primary ID contains a percent sign (e.g. `"50% done.pdf"`) could not be retrieved via POST. The ID is URL-decoded by the server, so the `%` was read as the start of an escape sequence; it is now escaped like every other string in the request body. Applies to both the sync and async clients. + +--- + ## [2.0.4] - 2026-05-18 ### Fixed @@ -74,7 +91,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - **`showSecrets()` removed.** Use `getSecrets()` instead. - **`runInstalledQuery()` now auto-selects GET or POST** based on `params` type (`dict` → POST, `str` → GET). Passing a raw query string with `usePost=True` raises `TigerGraphException`. -- **MCP tools moved to [`pytigergraph-mcp`](https://github.com/tigergraph/pytigergraph-mcp).** Do not import from `pyTigerGraph.mcp` directly. +- **MCP tools moved to [`tigergraph-mcp`](https://github.com/tigergraph/tigergraph-mcp).** Do not import from `pyTigerGraph.mcp` directly. ### New Features diff --git a/README.md b/README.md index 21722849..e63a9433 100644 --- a/README.md +++ b/README.md @@ -164,8 +164,9 @@ with TigerGraphConnection(...) as conn: ### Asynchronous mode (`AsyncTigerGraphConnection`) -- Uses a single `aiohttp.ClientSession` with an unbounded connection pool shared across all concurrent coroutines — no GIL, no thread-scheduling overhead. +- Uses a single `aiohttp.ClientSession` per event loop, with an unbounded connection pool shared across all concurrent coroutines — no GIL, no thread-scheduling overhead. - Typically achieves higher QPS and lower tail latency than the threaded sync mode for I/O-bound workloads. +- A connection may be used from more than one event loop: reused across separate `asyncio.run()` calls, or shared by threads that each run their own loop. Each loop gets its own session and pool. Use `async with` (or `await conn.aclose()`) to release sockets when finished. ```python import asyncio diff --git a/pyTigerGraph/__init__.py b/pyTigerGraph/__init__.py index 189e9644..9d27ecf4 100644 --- a/pyTigerGraph/__init__.py +++ b/pyTigerGraph/__init__.py @@ -7,7 +7,7 @@ try: __version__ = _pkg_version("pyTigerGraph") except PackageNotFoundError: - __version__ = "2.0.4" + __version__ = "2.0.5" __license__ = "Apache 2" diff --git a/pyTigerGraph/ai/ai.py b/pyTigerGraph/ai/ai.py index 7f5f339d..f4a1fe1c 100644 --- a/pyTigerGraph/ai/ai.py +++ b/pyTigerGraph/ai/ai.py @@ -201,17 +201,36 @@ def retrieveDocs(self, query: str, top_k: int = 3): "/retrieve_docs?top_k="+str(top_k) return self.conn._req("POST", url, authMode="pwd", data=data, jsonData=True, resKey=None, skipCheck=True) - def query(self, query): + def query(self, query, mode: str = None, rag_method: str = None, include_fields: list = None): """ Query the database with natural language. Args: query (str): Natural language query to ask about the database. + mode (str): + Chat engine to use: "agentic", "classic", or None to defer + to the graph's configured default. + rag_method (str): + Engine variant. When agentic: "auto", "planned", or + "reactive". When classic: "auto" or a retriever name + (e.g. "hybrid", "similarity", "contextual", + "entityrelationship", "community"). None defers to the + configured default. + include_fields (list): + Extra response fields beyond the answer. None returns the + answer only; pass field names (e.g. ["query_sources"]) or + ["all"] to include the supporting sources / trace. Returns: JSON including the natural language response, a answered_question flag, and answer sources. """ data = { "query": query } + if mode is not None: + data["mode"] = mode + if rag_method is not None: + data["rag_method"] = rag_method + if include_fields is not None: + data["include_fields"] = include_fields url = self.nlqs_host+"/"+self.conn.graphname+"/query" return self.conn._req("POST", url, authMode="pwd", data=data, jsonData=True, resKey=None) diff --git a/pyTigerGraph/common/base.py b/pyTigerGraph/common/base.py index e4e99674..bd1cd981 100644 --- a/pyTigerGraph/common/base.py +++ b/pyTigerGraph/common/base.py @@ -105,7 +105,7 @@ def __init__(self, host: str = "http://127.0.0.1", graphname: str = "", if inputHost.scheme not in ["http", "https"]: raise TigerGraphException("Invalid URL scheme. Supported schemes are http and https.", "E-0003") - # Extract port from URL if present (e.g. http://192.168.11.11:14240) + # Extract port from URL if present (e.g. http://127.0.0.1:14240) # Use hostname (without port) to avoid double-port URLs later. hostOnly = inputHost.hostname if not hostOnly: diff --git a/pyTigerGraph/common/query.py b/pyTigerGraph/common/query.py index 5631ad6f..3f7adbd4 100644 --- a/pyTigerGraph/common/query.py +++ b/pyTigerGraph/common/query.py @@ -69,7 +69,7 @@ def _parse_query_parameters(params: dict) -> str: f"Invalid vertex parameter '{k}': vertex type string must not be empty. " "Use (id,) for VERTEX or (id, 'type') for untyped VERTEX.") # VERTEX (untyped): (id, "type") → k=id&k.type=type - parts.append(k + "=" + str(v[0])) + parts.append(k + "=" + _safe_char(v[0])) parts.append(k + ".type=" + _safe_char(v[1])) else: raise TigerGraphException( @@ -88,7 +88,7 @@ def _parse_query_parameters(params: dict) -> str: "Use (id,) for VERTEX or (id, 'type') for untyped VERTEX.") # SET: (id, "type") → k[i]=id&k[i].type=type parts.append(k + "[" + str(i) + "]=" + _safe_char(vv[0])) - parts.append(k + "[" + str(i) + "].type=" + vv[1]) + parts.append(k + "[" + str(i) + "].type=" + _safe_char(vv[1])) else: raise TigerGraphException( f"Invalid vertex parameter '{k}[{i}]': expected (id,) for VERTEX " @@ -120,6 +120,20 @@ def _encode_str_for_post(value: str) -> str: return value.replace("%", "%25") +def _encode_vertex_id_for_post(value): + """Apply the same ``%`` escaping to a string vertex ID as to every other + string in the POST ``/query`` body. + + Vertex IDs go through the same URL-decoding as plain string parameters, so + a bare ``%`` in an ID is read as the start of a percent-escape and the ID + fails to resolve. Other reserved characters — spaces, ``/``, ``:``, ``@``, + ``+`` — and non-ASCII need no escaping and are left alone. + + Non-string IDs are returned unchanged so INT primary IDs keep their JSON type. + """ + return _encode_str_for_post(value) if isinstance(value, str) else value + + def _prep_query_parameters_json(params: dict) -> dict: """Converts a parameter dictionary into the JSON format expected by TigerGraph's POST /query endpoint. @@ -160,14 +174,14 @@ def _prep_query_parameters_json(params: dict) -> dict: if isinstance(v, tuple): if len(v) == 1: # VERTEX (typed): (id,) → {"id": id} - converted[k] = {"id": v[0]} + converted[k] = {"id": _encode_vertex_id_for_post(v[0])} elif len(v) == 2 and isinstance(v[1], str): if not v[1]: raise TigerGraphException( f"Invalid vertex parameter '{k}': vertex type string must not be empty. " "Use (id,) for VERTEX or (id, 'type') for untyped VERTEX.") # VERTEX (untyped): (id, "type") → {"id": id, "type": "type"} - converted[k] = {"id": v[0], "type": v[1]} + converted[k] = {"id": _encode_vertex_id_for_post(v[0]), "type": v[1]} else: raise TigerGraphException( f"Invalid vertex parameter '{k}': expected (id,) for VERTEX " @@ -188,14 +202,14 @@ def _prep_query_parameters_json(params: dict) -> dict: if isinstance(vv, tuple): if len(vv) == 1: # SET>: (id,) → {"id": id} - new_list.append({"id": vv[0]}) + new_list.append({"id": _encode_vertex_id_for_post(vv[0])}) elif len(vv) == 2 and isinstance(vv[1], str): if not vv[1]: raise TigerGraphException( f"Invalid vertex parameter '{k}': vertex type string must not be empty. " "Use (id,) for VERTEX or (id, 'type') for untyped VERTEX.") # SET: (id, "type") → {"id": id, "type": "type"} - new_list.append({"id": vv[0], "type": vv[1]}) + new_list.append({"id": _encode_vertex_id_for_post(vv[0]), "type": vv[1]}) else: raise TigerGraphException( f"Invalid vertex parameter '{k}': expected (id,) for VERTEX " diff --git a/pyTigerGraph/pyTigerGraphQuery.py b/pyTigerGraph/pyTigerGraphQuery.py index fa82d6ae..318fa6f9 100644 --- a/pyTigerGraph/pyTigerGraphQuery.py +++ b/pyTigerGraph/pyTigerGraphQuery.py @@ -365,6 +365,10 @@ def installQueries(self, queries: Union[str, list], flag: Union[str, list] = Non flag = ",".join(flag) params["flag"] = flag + # Install asynchronously so the server returns a requestId immediately + # instead of holding the request open for the whole compile. + params["async"] = "true" + res = self._req("GET", self.gsUrl + "/gsql/v1/queries/install", params=params, authMode="pwd", resKey=None) if wait: diff --git a/pyTigerGraph/pytgasync/pyTigerGraphBase.py b/pyTigerGraph/pytgasync/pyTigerGraphBase.py index af4e8c05..91c9167b 100644 --- a/pyTigerGraph/pytgasync/pyTigerGraphBase.py +++ b/pyTigerGraph/pytgasync/pyTigerGraphBase.py @@ -31,7 +31,7 @@ import logging import aiohttp -from typing import Optional, Union +from typing import Dict, Optional, Union from urllib.parse import urlparse from pyTigerGraph.common.auth import _is_auth_failure_response @@ -41,6 +41,25 @@ logger = logging.getLogger(__name__) + +class _LoopBinding: + """The HTTP session and locks bound to one event loop. + + aiohttp's ClientSession and asyncio.Lock both attach to the event loop that + is running when they are first used and cannot be moved to another, so they + are grouped here and looked up per loop. + """ + + __slots__ = ("client", "restpp_failover_lock", "token_refresh_lock") + + def __init__(self, client: aiohttp.ClientSession) -> None: + self.client = client + # Guards the one-time port failover (TG 3.x port 9000 → 4.x port 14240). + # Without it every concurrent task fails and enters the failover block at + # once, doubling requests and racing to overwrite restppUrl/restppPort. + self.restpp_failover_lock = asyncio.Lock() + self.token_refresh_lock = asyncio.Lock() + class AsyncPyTigerGraphBase(PyTigerGraphCore): def __init__(self, host: str = "http://127.0.0.1", graphname: str = "", gsqlSecret: str = "", username: str = "tigergraph", password: str = "tigergraph", @@ -101,15 +120,14 @@ def __init__(self, host: str = "http://127.0.0.1", graphname: str = "", version=version, apiToken=apiToken, useCert=useCert, certPath=certPath, debug=debug, sslPort=sslPort, gcp=gcp, jwtToken=jwtToken) - # Lazily initialized on first request (inside an async context) to avoid - # creating aiohttp.ClientSession outside an event loop in __init__. - self._async_client: Optional[aiohttp.ClientSession] = None - - # asyncio.Lock for the one-time port failover (TG 3.x port 9000 → 4.x port 14240). - # Without a lock all concurrent tasks simultaneously fail and all enter the failover - # block, doubling requests and racing to overwrite self.restppUrl/self.restppPort. - self._restpp_failover_lock = asyncio.Lock() - self._token_refresh_lock = asyncio.Lock() + # HTTP session and locks, one set per event loop, created on first use + # inside that loop. Both kinds of object are loop-bound and neither can + # move: aiohttp binds a ClientSession to the loop running when it is + # created, and an asyncio.Lock binds to the loop of its first await. + # Keying them by loop lets one connection serve sequential loops (the + # asyncio.run() pattern) and concurrent loops in separate threads + # without either interfering with the other. See _binding(). + self._loop_bindings: Dict[asyncio.AbstractEventLoop, _LoopBinding] = {} async def _req(self, method: str, url: str, authMode: str = "token", headers: dict = None, data: Union[dict, list, str] = None, resKey: str = "results", skipCheck: bool = False, @@ -144,8 +162,7 @@ async def _req(self, method: str, url: str, authMode: str = "token", headers: di The (relevant part of the) response from the request (as a dictionary). """ # Lazy init: session must be created inside an async context (event loop running). - if self._async_client is None or self._async_client.closed: - self._async_client = self._make_async_client() + binding = self._binding() _headers, _data, _ = self._prep_req(authMode, headers, url, method, data) @@ -179,7 +196,7 @@ async def _req(self, method: str, url: str, authMode: str = "token", headers: di if _is_auth_failure_response(_body): needs_token_retry = True if needs_token_retry: - async with self._token_refresh_lock: + async with binding.token_refresh_lock: if not getattr(self, "_refreshing_token", False): try: self._refreshing_token = True @@ -208,7 +225,7 @@ async def _req(self, method: str, url: str, authMode: str = "token", headers: di # ---- # Changes port to gsql port, adds /restpp to end to url, tries again, saves changes if successful if self.restppPort in url and "/gsql" not in url and ("/restpp" not in url or self.tgCloud): - async with self._restpp_failover_lock: + async with binding.restpp_failover_lock: if self.restppPort in url: newRestppUrl = self.host + ":" + self.gsPort + "/restpp" if "/restpp" in url: @@ -385,6 +402,74 @@ async def _delete(self, url: str, authMode: str = "token", headers: dict = None, return res + def _binding(self) -> "_LoopBinding": + """Return the HTTP session and locks belonging to the running event loop. + + Must be called from inside a coroutine. The binding is created on first + use in a loop and reused for that loop's lifetime, so concurrency within + one loop still shares a single connection pool. + + A connection outliving its loop is the normal result of driving it with + ``asyncio.run()``, which closes the loop when it returns. aiohttp does + not notice: ``session.closed`` stays False while the sockets underneath + are dead, and the next request fails with "RuntimeError: Event loop is + closed". Keying the session by loop avoids that, and — because two live + loops in separate threads get separate bindings — stops either thread + from tearing down the session the other is using. + + This method performs no ``await``, so it runs to completion without the + loop switching tasks. That is what makes first use safe for concurrent + callers on the same loop; do not introduce an await here. + """ + running = asyncio.get_running_loop() + + binding = self._loop_bindings.get(running) + if binding is not None and not binding.client.closed: + return binding + + self._prune_loop_bindings(running) + + binding = _LoopBinding(self._make_async_client()) + self._loop_bindings[running] = binding + return binding + + def _prune_loop_bindings(self, keep: Optional[asyncio.AbstractEventLoop] = None) -> None: + """Drop bindings whose event loop has been closed, releasing their sessions. + + Without this the registry would grow by one entry per ``asyncio.run()`` + call. The running loop is never pruned. + """ + for loop in [l for l in self._loop_bindings if l is not keep and l.is_closed()]: + binding = self._loop_bindings.pop(loop) + self._release_stale_client(binding.client, loop) + + @staticmethod + def _release_stale_client(client: Optional[aiohttp.ClientSession], + loop: Optional[asyncio.AbstractEventLoop]) -> None: + """Dispose of a session belonging to a loop that is no longer being used. + + When that loop is still running the session is closed on it properly. + When it is gone its sockets went with it, so the connector is only + marked closed — enough to keep aiohttp from reporting the session as + leaked at garbage-collection time. + """ + if client is None or client.closed: + return + + if loop is not None and not loop.is_closed() and loop.is_running(): + try: + loop.call_soon_threadsafe(loop.create_task, client.close()) + return + except RuntimeError: + pass # loop stopped between the check and the call + + connector = client.connector + if connector is not None: + try: + connector._close() + except Exception: # pragma: no cover - aiohttp internals moved + logger.debug("could not release stale HTTP connector", exc_info=True) + def _make_async_client(self) -> aiohttp.ClientSession: """Create a persistent aiohttp.ClientSession. @@ -426,7 +511,7 @@ async def _do_request( kwargs["json"] = _data else: kwargs["data"] = _data - async with self._async_client.request(method, url, **kwargs) as resp: + async with self._binding().client.request(method, url, **kwargs) as resp: # read() returns raw bytes — avoids charset detection overhead and lets # orjson/json.loads consume bytes directly without a decode step. body = await resp.read() @@ -443,28 +528,35 @@ async def aclose(self) -> None: await conn.runInstalledQuery(...) ``` """ - if self._async_client is not None and not self._async_client.closed: - await self._async_client.close() - self._async_client = None + bindings, self._loop_bindings = self._loop_bindings, {} + + try: + running = asyncio.get_running_loop() + except RuntimeError: # pragma: no cover - aclose is a coroutine + running = None + + for loop, binding in bindings.items(): + # Awaiting close() is only valid on the loop that owns the session; + # sessions belonging to any other loop go to the disposal path. + if loop is running and not binding.client.closed: + await binding.client.close() + else: + self._release_stale_client(binding.client, loop) def __del__(self) -> None: """Best-effort cleanup when the object is garbage-collected. - If the event loop is still running at GC time (e.g. during asyncio.run() - shutdown), schedules aclose() as a task so sockets are drained gracefully. - If the loop has already stopped, the OS reclaims the sockets and there is - nothing more we can do — this is not an error. + If the loop that owns the session is still running at GC time (e.g. during + asyncio.run() shutdown), the close is scheduled on it so sockets are + drained gracefully. If that loop has already stopped, the OS reclaims the + sockets and there is nothing more we can do — this is not an error. This does NOT replace explicit aclose() / async-with usage: GC timing is unpredictable and create_task() is fire-and-forget with no error handling. Use `async with AsyncTigerGraphConnection(...) as conn:` for reliable cleanup. """ - if self._async_client is not None and not self._async_client.closed: - try: - loop = asyncio.get_running_loop() - loop.create_task(self._async_client.close()) - except RuntimeError: - pass # no running loop; OS reclaims sockets on process exit + for loop, binding in (getattr(self, "_loop_bindings", None) or {}).items(): + self._release_stale_client(binding.client, loop) async def __aenter__(self): return self diff --git a/pyTigerGraph/pytgasync/pyTigerGraphQuery.py b/pyTigerGraph/pytgasync/pyTigerGraphQuery.py index 0291597f..c682ccb6 100644 --- a/pyTigerGraph/pytgasync/pyTigerGraphQuery.py +++ b/pyTigerGraph/pytgasync/pyTigerGraphQuery.py @@ -362,6 +362,10 @@ async def installQueries(self, queries: Union[str, list], flag: Union[str, list] flag = ",".join(flag) params["flag"] = flag + # Install asynchronously so the server returns a requestId immediately + # instead of holding the request open for the whole compile. + params["async"] = "true" + res = await self._req("GET", self.gsUrl + "/gsql/v1/queries/install", params=params, authMode="pwd", resKey=None) if wait: @@ -665,14 +669,18 @@ async def runInterpretedQuery(self, queryText: str, params: Union[str, dict] = N # - SET (no type): k[0]=id&k[0].type=vtype&k[1]=... if isinstance(params, dict): params = _parse_query_parameters(params) + # The query string is already percent-encoded, and aiohttp would encode + # it again if passed as params=, so append it to the URL ourselves (as + # _run_installed_query_get does). + query_string = "?" + str(params) if params else "" if await self._version_greater_than_4_0(): - ret = await self._req("POST", self.gsUrl + "/gsql/v1/queries/interpret", - params=params, data=queryText, authMode="pwd", + ret = await self._req("POST", self.gsUrl + "/gsql/v1/queries/interpret" + query_string, + data=queryText, authMode="pwd", headers={'Content-Type': 'text/plain'}) else: - ret = await self._req("POST", self.gsUrl + "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/gsqlserver/interpreted_query", data=queryText, - params=params, authMode="pwd") + ret = await self._req("POST", self.gsUrl + "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/gsqlserver/interpreted_query" + query_string, + data=queryText, authMode="pwd") if logger.level == logging.DEBUG: logger.debug("return: " + str(ret)) diff --git a/pyproject.toml b/pyproject.toml index d3ca2043..a71b8f15 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,19 +4,18 @@ build-backend = "setuptools.build_meta" [project] name = "pyTigerGraph" -version = "2.0.4" +version = "2.0.5" description = "Library to connect to TigerGraph databases" readme = "README.md" license = "Apache-2.0" authors = [{ name = "TigerGraph Inc.", email = "support@tigergraph.com" }] -requires-python = ">=3.8" +requires-python = ">=3.9" keywords = ["TigerGraph", "Graph Database", "Data Science", "Machine Learning"] classifiers = [ "Development Status :: 5 - Production/Stable", "Intended Audience :: Developers", "Topic :: Software Development :: Build Tools", "Programming Language :: Python :: 3", - "Programming Language :: Python :: 3.8", "Programming Language :: Python :: 3.9", "Programming Language :: Python :: 3.10", "Programming Language :: Python :: 3.11", diff --git a/pytigergraph-recipe/recipe/meta.yaml b/pytigergraph-recipe/recipe/meta.yaml index 70a29e60..e9263802 100644 --- a/pytigergraph-recipe/recipe/meta.yaml +++ b/pytigergraph-recipe/recipe/meta.yaml @@ -1,4 +1,4 @@ -{% set version = "2.0.4" %} +{% set version = "2.0.5" %} {% set python_min = "3.9" %} package: diff --git a/tests/README.md b/tests/README.md index 59fb5a64..d3ea8e0b 100644 --- a/tests/README.md +++ b/tests/README.md @@ -8,8 +8,9 @@ Most unit tests need an accessible TigerGraph database with a specific graph. No If you need to manually prepare a DB for testing, run the `testserver.gsql` script to create the graph for testing core functions (via the `gsql` command line tool; GraphStudio cannot be used.) The script will create a graph called "tests" and will populate it with various object types and some data. -⚠️ **NOTE**: The script drops all existing graphs and objects, so use it with a TigerGraph instance -that does not have operational or otherwise important data, schema design or code. +⚠️ **NOTE**: The script drops and recreates the "tests" graph, so do not run it against an instance +where a graph of that name holds anything worth keeping. Other graphs on the instance are left alone. +The script also creates the secrets `secret1`-`secret3` for the connecting user. About testing data for the GDS functions, please contact one of the maintainers. diff --git a/tests/pyTigerGraphUnitTest.py b/tests/pyTigerGraphUnitTest.py index cc59ceb8..d5f6b5d9 100644 --- a/tests/pyTigerGraphUnitTest.py +++ b/tests/pyTigerGraphUnitTest.py @@ -30,6 +30,17 @@ def make_connection(graphname: str = None): config = json.load(config_file) server_config.update(config) + # Environment variables override the file, so a live test instance can be + # pointed at without writing its address or credentials into a tracked file. + # Each key takes TG_TEST_, e.g. TG_TEST_HOST, TG_TEST_PASSWORD. + for key, default in list(server_config.items()) + [("getToken", False)]: + value = os.environ.get("TG_TEST_" + key.upper()) + if value is None: + continue + if isinstance(default, bool): + value = value.strip().lower() in ("1", "true", "yes", "on") + server_config[key] = value + conn = TigerGraphConnection( host=server_config["host"], graphname=graphname if graphname else server_config["graphname"], diff --git a/tests/pyTigerGraphUnitTestAsync.py b/tests/pyTigerGraphUnitTestAsync.py index a3af680a..8bb4628f 100644 --- a/tests/pyTigerGraphUnitTestAsync.py +++ b/tests/pyTigerGraphUnitTestAsync.py @@ -30,6 +30,17 @@ async def make_connection(graphname: str = None): config = json.load(config_file) server_config.update(config) + # Environment variables override the file, so a live test instance can be + # pointed at without writing its address or credentials into a tracked file. + # Each key takes TG_TEST_, e.g. TG_TEST_HOST, TG_TEST_PASSWORD. + for key, default in list(server_config.items()) + [("getToken", False)]: + value = os.environ.get("TG_TEST_" + key.upper()) + if value is None: + continue + if isinstance(default, bool): + value = value.strip().lower() in ("1", "true", "yes", "on") + server_config[key] = value + conn = AsyncTigerGraphConnection( host=server_config["host"], graphname=graphname if graphname else server_config["graphname"], diff --git a/tests/test_async_session_loop.py b/tests/test_async_session_loop.py new file mode 100644 index 00000000..dd357914 --- /dev/null +++ b/tests/test_async_session_loop.py @@ -0,0 +1,214 @@ +"""Regression tests for using an AsyncTigerGraphConnection across event loops. + +The HTTP session and the asyncio locks are bound to the event loop that created +them and cannot move. A connection must therefore cope with two shapes: + + * sequential loops -- asyncio.run() closes its loop on return, so the next + call runs on a new one (GML-2183); + * concurrent loops -- two threads each running their own loop against one + shared connection, which must not tear down each other's session. +""" + +import asyncio +import json +import threading +import unittest +import warnings +from http.server import BaseHTTPRequestHandler, HTTPServer + +import aiohttp + +from pyTigerGraph import AsyncTigerGraphConnection + + +class _Handler(BaseHTTPRequestHandler): + def do_GET(self): + body = json.dumps({"error": False, "results": "pong"}).encode() + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def log_message(self, *args): + pass + + +class _ServerCase(unittest.TestCase): + """Base case providing a local HTTP server and connection factory.""" + + def setUp(self): + self.server = HTTPServer(("127.0.0.1", 0), _Handler) + threading.Thread(target=self.server.serve_forever, daemon=True).start() + self.addCleanup(self.server.shutdown) + self.port = self.server.server_address[1] + self.url = f"http://127.0.0.1:{self.port}/echo" + + def conn(self): + # apiToken short-circuits token minting, so no TigerGraph is needed. + return AsyncTigerGraphConnection( + host="http://127.0.0.1", apiToken="dummy", + restppPort=str(self.port), gsPort="14240") + + async def _call(self, conn): + status, body, _ = await conn._do_request( + "GET", self.url, {}, None, False, None, + aiohttp.ClientTimeout(total=10)) + return status, json.loads(body)["results"] + + +class TestSequentialLoops(_ServerCase): + """A connection outliving its event loop, the asyncio.run() pattern.""" + + def test_request_succeeds_on_later_loop(self): + conn = self.conn() + with warnings.catch_warnings(): + warnings.simplefilter("error", ResourceWarning) + self.assertEqual(asyncio.run(self._call(conn)), (200, "pong")) + self.assertEqual(asyncio.run(self._call(conn)), (200, "pong")) + asyncio.run(conn.aclose()) + + def test_new_binding_per_loop(self): + conn = self.conn() + + async def touch(): + b = conn._binding() + return b, asyncio.get_running_loop() + + first, loop_a = asyncio.run(touch()) + second, loop_b = asyncio.run(touch()) + + self.assertIsNot(loop_a, loop_b) + self.assertIsNot(first, second) + self.assertIsNot(first.client, second.client) + self.assertIsNot(first.token_refresh_lock, second.token_refresh_lock) + + def test_binding_reused_within_one_loop(self): + conn = self.conn() + + async def twice(): + return conn._binding(), conn._binding() + + first, second = asyncio.run(twice()) + self.assertIs(first, second) + + def test_locks_usable_on_each_loop(self): + """A lock bound to a dead loop would raise when awaited on a new one.""" + conn = self.conn() + + async def acquire(): + b = conn._binding() + async with b.token_refresh_lock: + pass + async with b.restpp_failover_lock: + pass + + asyncio.run(acquire()) + asyncio.run(acquire()) # would raise "bound to a different event loop" + + def test_stale_session_not_reported_as_leaked(self): + conn = self.conn() + + async def touch(): + return conn._binding().client + + stale = asyncio.run(touch()) + self.assertFalse(stale.closed) # aiohttp does not notice the loop died + asyncio.run(touch()) + self.assertTrue(stale.closed) + + def test_closed_loops_pruned_from_registry(self): + """The registry must not grow by one entry per asyncio.run() call.""" + conn = self.conn() + + async def touch(): + conn._binding() + + for _ in range(5): + asyncio.run(touch()) + self.assertEqual(len(conn._loop_bindings), 1) + + +class TestConcurrentLoops(_ServerCase): + """Separate threads, each with its own loop, sharing one connection.""" + + def test_threads_do_not_close_each_others_session(self): + conn = self.conn() + results, errors = [], [] + + def worker(): + try: + for _ in range(8): + results.append(asyncio.run(self._call(conn))) + except Exception as e: # noqa: BLE001 - recorded and asserted below + errors.append(f"{type(e).__name__}: {e}") + + threads = [threading.Thread(target=worker) for _ in range(3)] + for t in threads: + t.start() + for t in threads: + t.join() + + self.assertEqual(errors, []) + self.assertEqual(results, [(200, "pong")] * 24) + + def test_concurrent_first_use_on_one_loop_shares_a_session(self): + """_binding() performs no await, so racing tasks cannot double-create.""" + conn = self.conn() + + async def many(): + out = await asyncio.gather(*(self._call(conn) for _ in range(20))) + return out, len(conn._loop_bindings) + + results, bindings = asyncio.run(many()) + self.assertEqual(results, [(200, "pong")] * 20) + self.assertEqual(bindings, 1) + + +class TestClose(_ServerCase): + """aclose() and __del__ across loop boundaries.""" + + def test_aclose_on_owning_loop(self): + conn = self.conn() + + async def run(): + await self._call(conn) + client = conn._binding().client + await conn.aclose() + return client + + client = asyncio.run(run()) + self.assertTrue(client.closed) + self.assertEqual(conn._loop_bindings, {}) + + def test_aclose_from_a_foreign_loop(self): + """Awaiting close() on the wrong loop would fail; it must be disposed instead.""" + conn = self.conn() + asyncio.run(self._call(conn)) + stale = next(iter(conn._loop_bindings.values())).client + + asyncio.run(conn.aclose()) + + self.assertTrue(stale.closed) + self.assertEqual(conn._loop_bindings, {}) + + def test_del_releases_every_binding(self): + conn = self.conn() + asyncio.run(self._call(conn)) + clients = [b.client for b in conn._loop_bindings.values()] + self.assertTrue(clients) + + conn.__del__() + + self.assertTrue(all(c.closed for c in clients)) + + def test_reusable_after_aclose(self): + conn = self.conn() + asyncio.run(self._call(conn)) + asyncio.run(conn.aclose()) + self.assertEqual(asyncio.run(self._call(conn)), (200, "pong")) + asyncio.run(conn.aclose()) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_common_query_helpers.py b/tests/test_common_query_helpers.py index f631d154..6abcc359 100644 --- a/tests/test_common_query_helpers.py +++ b/tests/test_common_query_helpers.py @@ -130,6 +130,36 @@ def test_tuple_untyped_vertex_integer_id(self): result = _prep_query_parameters_json({"v": (42, "Order")}) self.assertEqual(result, {"v": {"id": 42, "type": "Order"}}) + # ------------------------------------------------------------------ + # Vertex ID encoding — string IDs get the same % escaping as every other + # string in the POST body, because the server URL-decodes them too. + # ------------------------------------------------------------------ + + def test_tuple_vertex_id_percent_encoded(self): + # A bare % in a STRING id would otherwise be read as a percent-escape. + result = _prep_query_parameters_json({"doc": ("50% done.pdf",)}) + self.assertEqual(result, {"doc": {"id": "50%25 done.pdf"}}) + + def test_tuple_vertex_id_reserved_chars_untouched(self): + # Spaces, /, :, @, + and non-ASCII need no escaping — encoding them + # would be a no-op at best and is not what the server expects. + vid = "summer release_115605840 a/b:c@d+e \u65e5\u672c\u8a9e" + result = _prep_query_parameters_json({"v": (vid,)}) + self.assertEqual(result, {"v": {"id": vid}}) + + def test_tuple_vertex_id_int_not_encoded(self): + # INT primary IDs must keep their JSON type, not become a string. + result = _prep_query_parameters_json({"v": (42,)}) + self.assertEqual(result, {"v": {"id": 42}}) + + def test_untyped_vertex_id_percent_encoded(self): + result = _prep_query_parameters_json({"v": ("50% off", "Person")}) + self.assertEqual(result, {"v": {"id": "50%25 off", "type": "Person"}}) + + def test_set_vertex_ids_percent_encoded(self): + result = _prep_query_parameters_json({"vs": [("a%b",), ("c d",)]}) + self.assertEqual(result, {"vs": [{"id": "a%25b"}, {"id": "c d"}]}) + def test_tuple_invalid_3tuple_raises(self): with self.assertRaises(TigerGraphException): _prep_query_parameters_json({"v": ("id", "type", "extra")}) @@ -293,6 +323,23 @@ def test_untyped_vertex_set_list_of_2tuples(self): result = _parse_query_parameters({"vs": [("Tom", "Person"), ("Mary", "Person")]}) self.assertEqual(result, "vs[0]=Tom&vs[0].type=Person&vs[1]=Mary&vs[1].type=Person") + # ------------------------------------------------------------------ + # Vertex ID escaping — an unescaped & or # in an ID terminates the + # parameter (or starts a fragment) and the query string is truncated. + # ------------------------------------------------------------------ + + def test_untyped_vertex_2tuple_id_escaped(self): + result = _parse_query_parameters({"v": ("a&b#c d", "Person")}) + self.assertEqual(result, "v=a%26b%23c%20d&v.type=Person") + + def test_typed_vertex_1tuple_id_escaped(self): + result = _parse_query_parameters({"v": ("a&b#c d",)}) + self.assertEqual(result, "v=a%26b%23c%20d") + + def test_untyped_vertex_set_ids_and_types_escaped(self): + result = _parse_query_parameters({"vs": [("a&b", "Person")]}) + self.assertEqual(result, "vs[0]=a%26b&vs[0].type=Person") + def test_invalid_tuple_raises(self): with self.assertRaises(TigerGraphException): _parse_query_parameters({"v": ("id", "type", "extra")}) @@ -455,6 +502,22 @@ def test_db_typed_vertex_tuple_roundtrip(self): res = self._run({"p08_vertex_vertex4": (3, "vertex4")}) self.assertEqual(str(res[7]["p08_vertex_vertex4"]), "3") + def test_db_vertex_string_id_with_percent_roundtrip(self): + """A VERTEX whose STRING primary ID contains % must round-trip via POST. + The server URL-decodes vertex IDs, so an unescaped % makes the ID + unresolvable ('Failed to convert user vertex id').""" + vid = "50% done" + self.conn.upsertVertex("vertex1_all_types", vid, {}) + res = self._run({"p07_vertex": (vid, "vertex1_all_types")}) + self.assertEqual(str(res[6]["p07_vertex"]), vid) + + def test_db_vertex_string_id_reserved_chars_roundtrip(self): + """Spaces and other reserved characters are accepted as-is by POST.""" + vid = "summer release_115605840 a/b:c@d+e" + self.conn.upsertVertex("vertex1_all_types", vid, {}) + res = self._run({"p07_vertex": (vid, "vertex1_all_types")}) + self.assertEqual(str(res[6]["p07_vertex"]), vid) + # ------------------------------------------------------------------ # SET / BAG of scalars # ------------------------------------------------------------------ diff --git a/tests/test_interpreted_query_params.py b/tests/test_interpreted_query_params.py new file mode 100644 index 00000000..b7edb156 --- /dev/null +++ b/tests/test_interpreted_query_params.py @@ -0,0 +1,96 @@ +"""Regression tests for interpreted-query parameter encoding (GML-2310). + +Parameter values travel in the URL query string. They must be percent-encoded +exactly once, so the server's single decode yields the original value. The +async client used to pass the already-encoded string to aiohttp as params=, +which encoded it a second time. +""" + +import asyncio +import json +import threading +import unittest +from http.server import BaseHTTPRequestHandler, HTTPServer +from urllib.parse import parse_qs, urlsplit + +from pyTigerGraph import AsyncTigerGraphConnection, TigerGraphConnection + +VALUES = { + "plain": "abc", + "space": "a b", + "quotes": "it's \"x\"", + "amp": "x&y", + "eq": "k=v", + "pct": "50%", + "hash": "#1", + "plus": "1+1", + "unicode": "東京", +} + + +class _Handler(BaseHTTPRequestHandler): + def do_POST(self): + self.server.paths.append(self.path) + self.rfile.read(int(self.headers.get("Content-Length", 0))) + body = json.dumps({"error": False, "message": "", "results": [{}]}).encode() + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def log_message(self, *args): + pass + + +class TestInterpretedQueryParams(unittest.TestCase): + def setUp(self): + self.server = HTTPServer(("127.0.0.1", 0), _Handler) + self.server.paths = [] + threading.Thread(target=self.server.serve_forever, daemon=True).start() + self.addCleanup(self.server.shutdown) + self.kwargs = dict(host="http://127.0.0.1", graphname="g", apiToken="dummy", + restppPort=str(self.server.server_address[1]), + gsPort=str(self.server.server_address[1])) + + def _received(self): + """The parameters as the server sees them after decoding once.""" + query = urlsplit(self.server.paths[-1]).query + return {k: v[0] for k, v in parse_qs(query).items()} + + def _run_async(self, params, v4: bool): + async def run(): + conn = AsyncTigerGraphConnection(**self.kwargs) + + async def version_check(): + return v4 + conn._version_greater_than_4_0 = version_check + try: + await conn.runInterpretedQuery("INTERPRET QUERY () {}", params) + finally: + await conn.aclose() + asyncio.run(run()) + + def test_async_encodes_values_once(self): + for v4 in (True, False): + with self.subTest(v4=v4): + self._run_async(VALUES, v4) + self.assertEqual(self._received(), VALUES) + + def test_async_string_params_unchanged(self): + self._run_async("s=a%20b&n=1", True) + self.assertEqual(self._received(), {"s": "a b", "n": "1"}) + + def test_async_no_params(self): + self._run_async(None, True) + self.assertEqual(urlsplit(self.server.paths[-1]).query, "") + + def test_sync_matches_async(self): + conn = TigerGraphConnection(**self.kwargs) + conn._version_greater_than_4_0 = lambda: True + conn.runInterpretedQuery("INTERPRET QUERY () {}", VALUES) + self.assertEqual(self._received(), VALUES) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_v202_changes.py b/tests/test_v202_changes.py index 8d1e0743..d73d807e 100644 --- a/tests/test_v202_changes.py +++ b/tests/test_v202_changes.py @@ -855,8 +855,8 @@ class TestHostPortExtraction(unittest.TestCase): def test_port_in_url_sets_restpp_and_gs_port(self): """Port in URL should be used as restppPort and gsPort.""" - conn = _make_conn(host="http://192.168.11.11:14240") - self.assertEqual(conn.host, "http://192.168.11.11") + conn = _make_conn(host="http://127.0.0.1:14240") + self.assertEqual(conn.host, "http://127.0.0.1") self.assertEqual(conn.restppPort, "14240") self.assertEqual(conn.gsPort, "14240") self.assertIn("14240", conn.restppUrl) @@ -864,66 +864,66 @@ def test_port_in_url_sets_restpp_and_gs_port(self): def test_port_in_url_no_double_port(self): """URLs should never contain double ports.""" - conn = _make_conn(host="http://192.168.11.11:14240") + conn = _make_conn(host="http://127.0.0.1:14240") self.assertNotIn(":14240:14240", conn.restppUrl) self.assertNotIn(":14240:14240", conn.gsUrl) def test_no_port_in_url_uses_defaults(self): """Without port in URL, default ports should be used.""" - conn = _make_conn(host="http://192.168.11.11") - self.assertEqual(conn.host, "http://192.168.11.11") + conn = _make_conn(host="http://127.0.0.1") + self.assertEqual(conn.host, "http://127.0.0.1") self.assertEqual(conn.restppPort, "9000") self.assertEqual(conn.gsPort, "14240") def test_port_in_url_with_matching_restpp_port(self): """Explicit restppPort matching URL port should work.""" - conn = _make_conn(host="http://192.168.11.11:14240", restppPort="14240") + conn = _make_conn(host="http://127.0.0.1:14240", restppPort="14240") self.assertEqual(conn.restppPort, "14240") def test_port_in_url_conflicts_with_restpp_port(self): """Explicit non-default restppPort differing from URL port should raise.""" with self.assertRaises(TigerGraphException) as ctx: - _make_conn(host="http://192.168.11.11:14240", restppPort="7000") + _make_conn(host="http://127.0.0.1:14240", restppPort="7000") self.assertIn("conflicts", str(ctx.exception)) def test_port_in_url_with_only_gs_port_matching(self): """Explicit gsPort matching URL port: OK, URL port also sets restppPort.""" - conn = _make_conn(host="http://192.168.11.11:10000", gsPort="10000") + conn = _make_conn(host="http://127.0.0.1:10000", gsPort="10000") self.assertEqual(conn.restppPort, "10000") self.assertEqual(conn.gsPort, "10000") def test_port_in_url_conflicts_with_gs_port(self): """Explicit gsPort (only) differing from URL port: error.""" with self.assertRaises(TigerGraphException) as ctx: - _make_conn(host="http://192.168.11.11:7000", gsPort="10000") + _make_conn(host="http://127.0.0.1:7000", gsPort="10000") self.assertIn("conflicts", str(ctx.exception)) def test_both_ports_explicit_url_matches_restpp(self): """Both explicit, URL port matches restppPort: OK.""" - conn = _make_conn(host="http://192.168.11.11:7000", + conn = _make_conn(host="http://127.0.0.1:7000", restppPort="7000", gsPort="10000") - self.assertEqual(conn.host, "http://192.168.11.11") + self.assertEqual(conn.host, "http://127.0.0.1") self.assertEqual(conn.restppPort, "7000") self.assertEqual(conn.gsPort, "10000") def test_both_ports_explicit_url_matches_gs(self): """Both explicit, URL port matches gsPort: OK.""" - conn = _make_conn(host="http://192.168.11.11:10000", + conn = _make_conn(host="http://127.0.0.1:10000", restppPort="7000", gsPort="10000") - self.assertEqual(conn.host, "http://192.168.11.11") + self.assertEqual(conn.host, "http://127.0.0.1") self.assertEqual(conn.restppPort, "7000") self.assertEqual(conn.gsPort, "10000") def test_both_ports_explicit_url_matches_neither(self): """Both explicit, URL port matches neither: error.""" with self.assertRaises(TigerGraphException) as ctx: - _make_conn(host="http://192.168.11.11:5000", + _make_conn(host="http://127.0.0.1:5000", restppPort="7000", gsPort="10000") self.assertIn("conflicts", str(ctx.exception)) def test_port_in_url_with_matching_gs_port(self): """Explicit gsPort matching URL port should work.""" - conn = _make_conn(host="http://192.168.11.11:14240", gsPort="14240") + conn = _make_conn(host="http://127.0.0.1:14240", gsPort="14240") self.assertEqual(conn.gsPort, "14240") def test_https_with_port(self):