diff --git a/src/py/client/__init__.py b/src/py/client/__init__.py index d743c566..75ced954 100644 --- a/src/py/client/__init__.py +++ b/src/py/client/__init__.py @@ -1,7 +1,7 @@ # pyright: strict, reportUnknownMemberType=false, reportUnknownVariableType=false from __future__ import annotations -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any, Optional, Self import requests @@ -123,7 +123,7 @@ def __init__( else: raise - def __enter__(self, *_: Any) -> Client: + def __enter__(self, *_: Any) -> Self: return self def __exit__(self, *_: Any) -> None: @@ -244,11 +244,14 @@ def test_src( seed: Optional[int] = None, timeout: Optional[float] = None, ) -> simple_api_pb2.TestRes: - seed = seed or 0 + req = simple_api_pb2.TestSrcReq(src=src, session=self._sesh) + if seed is not None: + req.seed = seed + timeout = timeout or self._timeout return self._client.test_src( ctx=self.mk_context(), - request=simple_api_pb2.TestSrcReq(src=src, session=self._sesh, seed=seed), + request=req, timeout=timeout, ) @@ -266,13 +269,14 @@ def test_name( seed: Optional[int] = None, timeout: Optional[float] = None, ) -> simple_api_pb2.TestRes: - seed = seed or 0 + req = simple_api_pb2.TestNameReq(name=name, session=self._sesh) + if seed is not None: + req.seed = seed + timeout = timeout or self._timeout return self._client.test_name( ctx=self.mk_context(), - request=simple_api_pb2.TestNameReq( - name=name, session=self._sesh, seed=seed - ), + request=req, timeout=timeout, ) diff --git a/src/py/client/_async.py b/src/py/client/_async.py index 898d3616..f986a948 100644 --- a/src/py/client/_async.py +++ b/src/py/client/_async.py @@ -1,7 +1,7 @@ # pyright: strict, reportUnknownMemberType=false, reportUnknownVariableType=false from __future__ import annotations -from typing import Any, Optional +from typing import Any, Optional, Self import aiohttp # type: ignore[import-not-found] @@ -75,7 +75,7 @@ def __init__( ) self._timeout = timeout - async def __aenter__(self, *_: Any) -> AsyncClient: + async def __aenter__(self, *_: Any) -> Self: await self._session.__aenter__() if self._session_id is None: try: @@ -217,36 +217,56 @@ async def instance_src( timeout=timeout, ) - async def qcheck_src( + async def test_src( self, src: str, seed: Optional[int] = None, timeout: Optional[float] = None, ) -> simple_api_pb2.TestRes: - seed = seed or 0 + req = simple_api_pb2.TestSrcReq(src=src, session=self._sesh) + if seed is not None: + req.seed = seed + timeout = timeout or self._timeout - return await self._client.qcheck_src( + return await self._client.test_src( ctx=self.mk_context(), - request=simple_api_pb2.QCheckSrcReq(src=src, session=self._sesh, seed=seed), + request=req, timeout=timeout, ) - async def qcheck_name( + async def qcheck_src( + self, + src: str, + seed: Optional[int] = None, + timeout: Optional[float] = None, + ) -> simple_api_pb2.TestRes: + return await self.test_src(src=src, seed=seed, timeout=timeout) + + async def test_name( self, name: str, seed: Optional[int] = None, timeout: Optional[float] = None, ) -> simple_api_pb2.TestRes: - seed = seed or 0 + req = simple_api_pb2.TestNameReq(name=name, session=self._sesh) + if seed is not None: + req.seed = seed + timeout = timeout or self._timeout - return await self._client.qcheck_name( + return await self._client.test_name( ctx=self.mk_context(), - request=simple_api_pb2.QCheckNameReq( - name=name, session=self._sesh, seed=seed - ), + request=req, timeout=timeout, ) + async def qcheck_name( + self, + name: str, + seed: Optional[int] = None, + timeout: Optional[float] = None, + ) -> simple_api_pb2.TestRes: + return await self.test_name(name=name, seed=seed, timeout=timeout) + async def list_artifacts( self, task: task_pb2.Task, timeout: Optional[float] = None ) -> api_pb2.ArtifactListResult: diff --git a/src/py/client/_common.py b/src/py/client/_common.py index 453a7c91..d53b2c32 100644 --- a/src/py/client/_common.py +++ b/src/py/client/_common.py @@ -21,4 +21,8 @@ def is_session_not_found(ex: TwirpServerException) -> bool: # Replace with a typed code check once the server is updated. body = (getattr(ex, "meta", None) or {}).get("body") or {} # type: ignore msg = body.get("msg") or "" # type: ignore - return "Session not found" in msg or "Unknown session" in msg + return ( + "Session not found" in msg + or "Unknown session" in msg + or "InvalidSession" in msg + ) diff --git a/src/py/pyproject.toml b/src/py/pyproject.toml index 563bc4d6..26d51276 100644 --- a/src/py/pyproject.toml +++ b/src/py/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "imandrax_api" -version = "0.20.1" +version = "0.20.2" description = "Imandrax API client library" requires-python = ">=3.12" dependencies = [ diff --git a/src/py/setup.py b/src/py/setup.py index a05f6ae6..6b70cb86 100644 --- a/src/py/setup.py +++ b/src/py/setup.py @@ -1,6 +1,6 @@ from setuptools import setup -VERSION = "0.20.1" +VERSION = "0.20.2" setup( name="imandrax_api", version=VERSION, diff --git a/src/py/uv.lock b/src/py/uv.lock index 73f63b44..35e10436 100644 --- a/src/py/uv.lock +++ b/src/py/uv.lock @@ -284,7 +284,7 @@ wheels = [ [[package]] name = "imandrax-api" -version = "0.20.1" +version = "0.20.2" source = { editable = "." } dependencies = [ { name = "protobuf" },