Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 28 additions & 2 deletions src/postgrest/run-unasync.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,11 +2,37 @@

import unasync

rules = (
unasync.Rule(
fromdir="/_async/",
todir="/_sync/",
additional_replacements={
"AsyncClient": "Client",
"AsyncPostgrestClient": "SyncPostgrestClient",
"async_postgrest_client": "sync_postgrest_client",
"AsyncFilterRequestBuilder": "SyncFilterRequestBuilder",
"AsyncSelectRequestBuilder": "SyncSelectRequestBuilder",
"AsyncQueryRequestBuilder": "SyncQueryRequestBuilder",
"AsyncSingleRequestBuilder": "SyncSingleRequestBuilder",
"AsyncMaybeSingleRequestBuilder": "SyncMaybeSingleRequestBuilder",
"AsyncExplainRequestBuilder": "SyncExplainRequestBuilder",
"AsyncRPCFilterRequestBuilder": "SyncRPCFilterRequestBuilder",
"AsyncRequestBuilder": "SyncRequestBuilder",
"AsyncHTTPTransport": "HTTPTransport",
"aclose": "close",
"asyncio.sleep": "time.sleep",
"import asyncio": "import time",
"postgrest._async.request_builder.AsyncSelectRequestBuilder.execute": "postgrest._sync.request_builder.SyncSelectRequestBuilder.execute",
"httpx._client.AsyncClient.request": "httpx._client.Client.request",
"@pytest.mark.asyncio\n": "",
" @pytest.mark.asyncio\n": "",
},
),
unasync._DEFAULT_RULE,
)
paths = Path("src/postgrest").glob("**/*.py")
tests = Path("tests").glob("**/*.py")

rules = (unasync._DEFAULT_RULE,)

files = [str(p) for p in list(paths) + list(tests)]

if __name__ == "__main__":
Expand Down
25 changes: 25 additions & 0 deletions src/postgrest/src/postgrest/_async/request_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,10 @@ class AsyncQueryRequestBuilder:
def __init__(self, request: ReqConfig):
self.request = request

def clone(self: Self) -> Self:
"""Create a clone of this query request builder."""
return self.__class__(self.request.clone())

def select(self: QueryBuilderT, *columns: str) -> QueryBuilderT:
_, params, _, _ = pre_select(*columns, count=None)
self.request.params = self.request.params.add("select", params["select"])
Expand Down Expand Up @@ -102,6 +106,10 @@ class AsyncSingleRequestBuilder:
def __init__(self, request: ReqConfig):
self.request = request

def clone(self: Self) -> Self:
"""Create a clone of this single request builder."""
return self.__class__(self.request.clone())

def retry(self, enabled: bool) -> Self:
self.request.retry_enabled = enabled
return self
Expand Down Expand Up @@ -135,6 +143,10 @@ class AsyncExplainRequestBuilder:
def __init__(self, request: ReqConfig):
self.request = request

def clone(self: Self) -> Self:
"""Create a clone of this explain request builder."""
return self.__class__(self.request.clone())

def retry(self, enabled: bool) -> Self:
self.request.retry_enabled = enabled
return self
Expand All @@ -155,6 +167,10 @@ class AsyncMaybeSingleRequestBuilder:
def __init__(self, request: ReqConfig):
self.request = request

def clone(self: Self) -> Self:
"""Create a clone of this maybe single request builder."""
return self.__class__(self.request.clone())

def retry(self, enabled: bool) -> Self:
self.request.retry_enabled = enabled
return self
Expand Down Expand Up @@ -301,6 +317,15 @@ def __init__(
self.headers = headers
self.auth = auth

def clone(self: Self) -> Self:
"""Create a clone of this request builder."""
return self.__class__(
session=self.session,
path=self.path,
headers=Headers(self.headers),
auth=self.auth,
)

def select(
self,
*columns: str,
Expand Down
6 changes: 3 additions & 3 deletions src/postgrest/src/postgrest/_sync/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,9 +17,9 @@
from ..types import CountMethod
from ..version import __version__
from .request_builder import (
RequestConfig,
SyncRequestBuilder,
SyncRPCFilterRequestBuilder,
RequestConfig,
)


Expand Down Expand Up @@ -119,9 +119,9 @@ def __enter__(self) -> SyncPostgrestClient:
return self

def __exit__(self, exc_type, exc, tb) -> None:
self.aclose()
self.close()

def aclose(self) -> None:
def close(self) -> None:
"""Close the underlying HTTP connections."""
self.session.close()

Expand Down
41 changes: 33 additions & 8 deletions src/postgrest/src/postgrest/_sync/request_builder.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,9 @@
from __future__ import annotations

import time
import asyncio
from typing import Any, Generic, Literal, Optional, TypeVar, Union, overload

from httpx import BasicAuth, Client, Headers, QueryParams, Response
from httpx import Client, BasicAuth, Headers, QueryParams, Response
from pydantic import ValidationError
from typing_extensions import Self, override
from yarl import URL
Expand Down Expand Up @@ -51,7 +51,7 @@ def send_with_retry(req: ReqConfig) -> Response:
resp = req.send(headers)
if resp.is_success or not req.should_retry(resp, attempt_count=attempt_count):
break
time.sleep(get_retry_delay(resp, attempt_count))
asyncio.sleep(get_retry_delay(resp, attempt_count))
attempt_count += 1
return resp

Expand All @@ -60,6 +60,10 @@ class SyncQueryRequestBuilder:
def __init__(self, request: ReqConfig):
self.request = request

def clone(self: Self) -> Self:
"""Create a clone of this query request builder."""
return self.__class__(self.request.clone())

def select(self: QueryBuilderT, *columns: str) -> QueryBuilderT:
_, params, _, _ = pre_select(*columns, count=None)
self.request.params = self.request.params.add("select", params["select"])
Expand Down Expand Up @@ -102,6 +106,10 @@ class SyncSingleRequestBuilder:
def __init__(self, request: ReqConfig):
self.request = request

def clone(self: Self) -> Self:
"""Create a clone of this single request builder."""
return self.__class__(self.request.clone())

def retry(self, enabled: bool) -> Self:
self.request.retry_enabled = enabled
return self
Expand Down Expand Up @@ -135,6 +143,10 @@ class SyncExplainRequestBuilder:
def __init__(self, request: ReqConfig):
self.request = request

def clone(self: Self) -> Self:
"""Create a clone of this explain request builder."""
return self.__class__(self.request.clone())

def retry(self, enabled: bool) -> Self:
self.request.retry_enabled = enabled
return self
Expand All @@ -155,6 +167,10 @@ class SyncMaybeSingleRequestBuilder:
def __init__(self, request: ReqConfig):
self.request = request

def clone(self: Self) -> Self:
"""Create a clone of this maybe single request builder."""
return self.__class__(self.request.clone())

def retry(self, enabled: bool) -> Self:
self.request.retry_enabled = enabled
return self
Expand Down Expand Up @@ -301,6 +317,15 @@ def __init__(
self.headers = headers
self.auth = auth

def clone(self: Self) -> Self:
"""Create a clone of this request builder."""
return self.__class__(
session=self.session,
path=self.path,
headers=Headers(self.headers),
auth=self.auth,
)

def select(
self,
*columns: str,
Expand All @@ -313,7 +338,7 @@ def select(
*columns: The names of the columns to fetch.
count: The method to use to get the count of rows returned.
Returns:
:class:`SyncSelectRequestBuilder`
:class:`AsyncSelectRequestBuilder`
"""
method, params, headers, json = pre_select(*columns, count=count, head=head)
headers.update(self.headers)
Expand Down Expand Up @@ -348,7 +373,7 @@ def insert(
Otherwise, use the default value for the column.
Only applies for bulk inserts.
Returns:
:class:`SyncQueryRequestBuilder`
:class:`AsyncQueryRequestBuilder`
"""
method, params, headers, json = pre_insert(
json,
Expand Down Expand Up @@ -392,7 +417,7 @@ def upsert(
not when merging with existing rows under `ignoreDuplicates: false`.
This also only applies when doing bulk upserts.
Returns:
:class:`SyncQueryRequestBuilder`
:class:`AsyncQueryRequestBuilder`
"""
method, params, headers, json = pre_upsert(
json,
Expand Down Expand Up @@ -428,7 +453,7 @@ def update(
count: The method to use to get the count of rows returned.
returning: Either 'minimal' or 'representation'
Returns:
:class:`SyncFilterRequestBuilder`
:class:`AsyncFilterRequestBuilder`
"""
method, params, headers, json = pre_update(
json,
Expand Down Expand Up @@ -459,7 +484,7 @@ def delete(
count: The method to use to get the count of rows returned.
returning: Either 'minimal' or 'representation'
Returns:
:class:`SyncFilterRequestBuilder`
:class:`AsyncFilterRequestBuilder`
"""
method, params, headers, json = pre_delete(
count=count,
Expand Down
20 changes: 20 additions & 0 deletions src/postgrest/src/postgrest/base_request_builder.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from __future__ import annotations

import copy
import json
import sys
from json import JSONDecodeError
Expand Down Expand Up @@ -76,6 +77,19 @@ def __init__(
self.auth = auth
self.retry_enabled = retry_enabled

def clone(self: Self) -> Self:
"""Create an independent clone of this request configuration."""
return self.__class__(
session=self.session,
path=self.path,
http_method=self.http_method,
headers=Headers(self.headers),
params=QueryParams(self.params),
auth=self.auth,
json=copy.deepcopy(self.json) if self.json is not None else None,
retry_enabled=self.retry_enabled,
)

@overload
def send(
self: RequestConfig[Client], additional_headers: Headers
Expand Down Expand Up @@ -282,6 +296,12 @@ def __init__(self, request: RequestConfig[C]) -> None:
self.request: RequestConfig[C] = request
self.negate_next = False

def clone(self: Self) -> Self:
"""Create a clone of this filter request builder with independent query parameters and headers."""
cloned = self.__class__(self.request.clone())
cloned.negate_next = self.negate_next
return cloned

@property
def not_(self: Self) -> Self:
"""Whether the filter applied next should be negated."""
Expand Down
33 changes: 33 additions & 0 deletions src/postgrest/tests/_async/test_filter_request_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -335,3 +335,36 @@ def test_max_affected_returns_self(filter_request_builder):
builder = filter_request_builder.max_affected(1)

assert builder is filter_request_builder


def test_filter_builder_clone_isolation(filter_request_builder: AsyncFilterRequestBuilder):
builder = filter_request_builder.eq("org_id", 42)
cloned = builder.clone()

cloned.eq("status", "active")
builder.eq("status", "archived")

assert str(builder.request.params) == "org_id=eq.42&status=eq.archived"
assert str(cloned.request.params) == "org_id=eq.42&status=eq.active"


def test_filter_builder_clone_negate_next(filter_request_builder: AsyncFilterRequestBuilder):
builder = filter_request_builder.not_
cloned = builder.clone()

assert cloned.negate_next is True

cloned.eq("name", "Alice")
assert str(cloned.request.params) == "name=not.eq.Alice"
assert cloned.negate_next is False
assert builder.negate_next is True


def test_filter_builder_clone_headers_isolation(filter_request_builder: AsyncFilterRequestBuilder):
builder = filter_request_builder.eq("id", 1)
cloned = builder.clone()

cloned.max_affected(5)

assert "prefer" in cloned.request.headers
assert "prefer" not in builder.request.headers
38 changes: 38 additions & 0 deletions src/postgrest/tests/_async/test_request_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -569,3 +569,41 @@ def test_single_with_csv_data(
)
assert isinstance(result.data, str)
assert result.data == csv_api_response


class TestClone:
def test_request_builder_clone(self, request_builder: AsyncRequestBuilder):
cloned = request_builder.clone()
cloned.headers["X-Custom"] = "test-value"

assert "X-Custom" in cloned.headers
assert "X-Custom" not in request_builder.headers

def test_select_builder_clone_branching(
self, request_builder: AsyncRequestBuilder
):
base = request_builder.select("*").eq("tenant_id", "t1")
branch_a = base.clone().order("created_at", desc=True).limit(10)
branch_b = base.clone().order("name").range(20, 30)

assert str(base.request.params) == "select=%2A&tenant_id=eq.t1"
assert (
str(branch_a.request.params)
== "select=%2A&tenant_id=eq.t1&order=created_at.desc&limit=10"
)
assert (
str(branch_b.request.params)
== "select=%2A&tenant_id=eq.t1&order=name.asc&offset=20&limit=11"
)

def test_single_builder_clone(self, request_builder: AsyncRequestBuilder):
single = request_builder.select("*").single()
cloned = single.clone()

cloned.request.headers["X-Single-Test"] = "val"

assert (
cloned.request.headers["Accept"] == "application/vnd.pgrst.object+json"
)
assert "X-Single-Test" in cloned.request.headers
assert "X-Single-Test" not in single.request.headers
Loading