dqlite-client 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,78 @@
1
+ Metadata-Version: 2.4
2
+ Name: dqlite-client
3
+ Version: 0.1.0
4
+ Summary: Async Python client for dqlite with connection pooling and leader detection
5
+ Project-URL: Homepage, https://github.com/letsdiscodev/python-dqlite-client
6
+ Project-URL: Repository, https://github.com/letsdiscodev/python-dqlite-client
7
+ Project-URL: Issues, https://github.com/letsdiscodev/python-dqlite-client/issues
8
+ Author-email: Antoine Leclair <antoineleclair@gmail.com>
9
+ License-Expression: MIT
10
+ Keywords: async,asyncio,database,distributed,dqlite,sqlite
11
+ Classifier: Development Status :: 3 - Alpha
12
+ Classifier: Framework :: AsyncIO
13
+ Classifier: Intended Audience :: Developers
14
+ Classifier: License :: OSI Approved :: MIT License
15
+ Classifier: Operating System :: OS Independent
16
+ Classifier: Programming Language :: Python :: 3
17
+ Classifier: Programming Language :: Python :: 3.13
18
+ Classifier: Topic :: Database
19
+ Classifier: Topic :: Database :: Database Engines/Servers
20
+ Classifier: Typing :: Typed
21
+ Requires-Python: >=3.13
22
+ Requires-Dist: dqlite-wire>=0.1.0
23
+ Provides-Extra: dev
24
+ Requires-Dist: mypy>=1.0; extra == 'dev'
25
+ Requires-Dist: pytest-asyncio>=0.23; extra == 'dev'
26
+ Requires-Dist: pytest-cov>=4.0; extra == 'dev'
27
+ Requires-Dist: pytest>=8.0; extra == 'dev'
28
+ Requires-Dist: ruff>=0.4; extra == 'dev'
29
+ Provides-Extra: yaml
30
+ Requires-Dist: pyyaml>=6.0; extra == 'yaml'
31
+ Description-Content-Type: text/markdown
32
+
33
+ # dqlite-client
34
+
35
+ Async Python client for [dqlite](https://dqlite.io/), following asyncpg patterns.
36
+
37
+ ## Installation
38
+
39
+ ```bash
40
+ pip install dqlite-client
41
+ ```
42
+
43
+ ## Usage
44
+
45
+ ```python
46
+ import asyncio
47
+ from dqliteclient import connect
48
+
49
+ async def main():
50
+ conn = await connect("localhost:9001")
51
+ async with conn.transaction():
52
+ await conn.execute("CREATE TABLE IF NOT EXISTS test (id INTEGER PRIMARY KEY, name TEXT)")
53
+ await conn.execute("INSERT INTO test (name) VALUES (?)", ["hello"])
54
+ rows = await conn.fetch("SELECT * FROM test")
55
+ for row in rows:
56
+ print(row)
57
+ await conn.close()
58
+
59
+ asyncio.run(main())
60
+ ```
61
+
62
+ ## Connection Pooling
63
+
64
+ ```python
65
+ from dqliteclient import create_pool
66
+
67
+ pool = await create_pool(["localhost:9001", "localhost:9002", "localhost:9003"])
68
+ async with pool.acquire() as conn:
69
+ rows = await conn.fetch("SELECT 1")
70
+ ```
71
+
72
+ ## Development
73
+
74
+ See [DEVELOPMENT.md](DEVELOPMENT.md) for setup and contribution guidelines.
75
+
76
+ ## License
77
+
78
+ MIT
@@ -0,0 +1,12 @@
1
+ dqliteclient/__init__.py,sha256=28ef1wvsEvxLr8Q8qQ4PbpQGdy5pq6AzvgdrLuxjZHU,1934
2
+ dqliteclient/cluster.py,sha256=44FCWwnFsvMrPZetoYShES6Jsza_QRHZoZ468xS_LE4,3361
3
+ dqliteclient/connection.py,sha256=-dIk_BS8WOUKpm5XXBBpEj9XuiVdITZOy5kQl_GsiRU,4971
4
+ dqliteclient/exceptions.py,sha256=eq1t7WEyWg33N4rGaEyW69kouduIjs7bcEThhlGPvZo,682
5
+ dqliteclient/node_store.py,sha256=OnhTneRakDeMrvzCTdodtmGLvXsAQZd8BEO1rdu3io8,1236
6
+ dqliteclient/pool.py,sha256=2TJ5PRy0TJky3LzlE0MNEksH8jyMdtWt9q6AM5D5GRw,4125
7
+ dqliteclient/protocol.py,sha256=0iIlkdhIrQtvrMF0DJFTyPr_qLnV-qUaNdFJnWUuaFw,6991
8
+ dqliteclient/py.typed,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
9
+ dqliteclient/retry.py,sha256=wYtat__HP5jc3MWcimP0BNUqZfGc1FiM3qDrpTRpxqs,1327
10
+ dqlite_client-0.1.0.dist-info/METADATA,sha256=t2LfkDweHFimfNCKfkoahG1MhVunOeBqlBba18DSH5E,2346
11
+ dqlite_client-0.1.0.dist-info/WHEEL,sha256=WLgqFyCfm_KASv4WHyYy0P3pM_m7J5L9k2skdKLirC8,87
12
+ dqlite_client-0.1.0.dist-info/RECORD,,
@@ -0,0 +1,4 @@
1
+ Wheel-Version: 1.0
2
+ Generator: hatchling 1.28.0
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any
@@ -0,0 +1,82 @@
1
+ """Async Python client for dqlite."""
2
+
3
+ from dqliteclient.cluster import ClusterClient
4
+ from dqliteclient.connection import DqliteConnection
5
+ from dqliteclient.exceptions import (
6
+ ClusterError,
7
+ ConnectionError,
8
+ DqliteError,
9
+ OperationalError,
10
+ ProtocolError,
11
+ )
12
+ from dqliteclient.node_store import MemoryNodeStore, NodeStore
13
+ from dqliteclient.pool import ConnectionPool
14
+
15
+ __all__ = [
16
+ "connect",
17
+ "create_pool",
18
+ "DqliteConnection",
19
+ "ConnectionPool",
20
+ "ClusterClient",
21
+ "NodeStore",
22
+ "MemoryNodeStore",
23
+ "DqliteError",
24
+ "ConnectionError",
25
+ "ProtocolError",
26
+ "ClusterError",
27
+ "OperationalError",
28
+ ]
29
+
30
+ __version__ = "0.1.0"
31
+
32
+
33
+ async def connect(
34
+ address: str,
35
+ *,
36
+ database: str = "default",
37
+ timeout: float = 10.0,
38
+ ) -> DqliteConnection:
39
+ """Connect to a dqlite node.
40
+
41
+ Args:
42
+ address: Node address in "host:port" format
43
+ database: Database name to open
44
+ timeout: Connection timeout in seconds
45
+
46
+ Returns:
47
+ A connected DqliteConnection
48
+ """
49
+ conn = DqliteConnection(address, database=database, timeout=timeout)
50
+ await conn.connect()
51
+ return conn
52
+
53
+
54
+ async def create_pool(
55
+ addresses: list[str],
56
+ *,
57
+ database: str = "default",
58
+ min_size: int = 1,
59
+ max_size: int = 10,
60
+ timeout: float = 10.0,
61
+ ) -> ConnectionPool:
62
+ """Create a connection pool with automatic leader detection.
63
+
64
+ Args:
65
+ addresses: List of node addresses in "host:port" format
66
+ database: Database name to open
67
+ min_size: Minimum number of connections to maintain
68
+ max_size: Maximum number of connections
69
+ timeout: Connection timeout in seconds
70
+
71
+ Returns:
72
+ An initialized ConnectionPool
73
+ """
74
+ pool = ConnectionPool(
75
+ addresses,
76
+ database=database,
77
+ min_size=min_size,
78
+ max_size=max_size,
79
+ timeout=timeout,
80
+ )
81
+ await pool.initialize()
82
+ return pool
@@ -0,0 +1,105 @@
1
+ """Cluster management and leader detection for dqlite."""
2
+
3
+ import asyncio
4
+
5
+ from dqliteclient.connection import DqliteConnection
6
+ from dqliteclient.exceptions import ClusterError
7
+ from dqliteclient.node_store import MemoryNodeStore, NodeInfo, NodeStore
8
+ from dqliteclient.protocol import DqliteProtocol
9
+ from dqliteclient.retry import retry_with_backoff
10
+
11
+
12
+ class ClusterClient:
13
+ """Client with automatic leader detection and failover."""
14
+
15
+ def __init__(
16
+ self,
17
+ node_store: NodeStore,
18
+ *,
19
+ timeout: float = 10.0,
20
+ ) -> None:
21
+ """Initialize cluster client.
22
+
23
+ Args:
24
+ node_store: Store for cluster node information
25
+ timeout: Connection timeout in seconds
26
+ """
27
+ self._node_store = node_store
28
+ self._timeout = timeout
29
+ self._leader_address: str | None = None
30
+
31
+ @classmethod
32
+ def from_addresses(cls, addresses: list[str], timeout: float = 10.0) -> "ClusterClient":
33
+ """Create cluster client from list of addresses."""
34
+ store = MemoryNodeStore(addresses)
35
+ return cls(store, timeout=timeout)
36
+
37
+ async def find_leader(self) -> str:
38
+ """Find the current cluster leader.
39
+
40
+ Returns the leader address.
41
+ """
42
+ nodes = await self._node_store.get_nodes()
43
+
44
+ if not nodes:
45
+ raise ClusterError("No nodes configured")
46
+
47
+ errors: list[str] = []
48
+
49
+ for node in nodes:
50
+ try:
51
+ leader_address = await self._query_leader(node.address)
52
+ if leader_address:
53
+ self._leader_address = leader_address
54
+ return leader_address
55
+ except Exception as e:
56
+ errors.append(f"{node.address}: {e}")
57
+ continue
58
+
59
+ raise ClusterError(f"Could not find leader. Errors: {'; '.join(errors)}")
60
+
61
+ async def _query_leader(self, address: str) -> str | None:
62
+ """Query a node for the current leader."""
63
+ host, port_str = address.rsplit(":", 1)
64
+ port = int(port_str)
65
+
66
+ try:
67
+ reader, writer = await asyncio.wait_for(
68
+ asyncio.open_connection(host, port),
69
+ timeout=self._timeout,
70
+ )
71
+ except (TimeoutError, OSError):
72
+ return None
73
+
74
+ protocol = DqliteProtocol(reader, writer)
75
+
76
+ try:
77
+ await protocol.handshake()
78
+ node_id, leader_addr = await protocol.get_leader()
79
+
80
+ # If address is empty, this node is the leader
81
+ if not leader_addr:
82
+ return address
83
+
84
+ return leader_addr
85
+ finally:
86
+ protocol.close()
87
+ await protocol.wait_closed()
88
+
89
+ async def connect(self, database: str = "default") -> DqliteConnection:
90
+ """Connect to the cluster leader.
91
+
92
+ Returns a connection to the current leader.
93
+ """
94
+
95
+ async def try_connect() -> DqliteConnection:
96
+ leader = await self.find_leader()
97
+ conn = DqliteConnection(leader, database=database, timeout=self._timeout)
98
+ await conn.connect()
99
+ return conn
100
+
101
+ return await retry_with_backoff(try_connect, max_attempts=5)
102
+
103
+ async def update_nodes(self, nodes: list[NodeInfo]) -> None:
104
+ """Update the node store with new node information."""
105
+ await self._node_store.set_nodes(nodes)
@@ -0,0 +1,146 @@
1
+ """High-level connection interface for dqlite."""
2
+
3
+ import asyncio
4
+ from collections.abc import AsyncIterator
5
+ from contextlib import asynccontextmanager
6
+ from typing import Any
7
+
8
+ from dqliteclient.exceptions import ConnectionError
9
+ from dqliteclient.protocol import DqliteProtocol
10
+
11
+
12
+ class DqliteConnection:
13
+ """High-level async connection to a dqlite database."""
14
+
15
+ def __init__(
16
+ self,
17
+ address: str,
18
+ *,
19
+ database: str = "default",
20
+ timeout: float = 10.0,
21
+ ) -> None:
22
+ """Initialize connection (does not connect yet).
23
+
24
+ Args:
25
+ address: Node address in "host:port" format
26
+ database: Database name to open
27
+ timeout: Connection timeout in seconds
28
+ """
29
+ self._address = address
30
+ self._database = database
31
+ self._timeout = timeout
32
+ self._protocol: DqliteProtocol | None = None
33
+ self._db_id: int | None = None
34
+ self._in_transaction = False
35
+
36
+ @property
37
+ def address(self) -> str:
38
+ """Get the connection address."""
39
+ return self._address
40
+
41
+ @property
42
+ def is_connected(self) -> bool:
43
+ """Check if connected."""
44
+ return self._protocol is not None
45
+
46
+ async def connect(self) -> None:
47
+ """Establish connection to the database."""
48
+ if self._protocol is not None:
49
+ return
50
+
51
+ host, port_str = self._address.rsplit(":", 1)
52
+ port = int(port_str)
53
+
54
+ try:
55
+ reader, writer = await asyncio.wait_for(
56
+ asyncio.open_connection(host, port),
57
+ timeout=self._timeout,
58
+ )
59
+ except TimeoutError as e:
60
+ raise ConnectionError(f"Connection to {self._address} timed out") from e
61
+ except OSError as e:
62
+ raise ConnectionError(f"Failed to connect to {self._address}: {e}") from e
63
+
64
+ self._protocol = DqliteProtocol(reader, writer)
65
+
66
+ try:
67
+ await self._protocol.handshake()
68
+ self._db_id = await self._protocol.open_database(self._database)
69
+ except Exception:
70
+ self._protocol.close()
71
+ self._protocol = None
72
+ raise
73
+
74
+ async def close(self) -> None:
75
+ """Close the connection."""
76
+ if self._protocol is not None:
77
+ self._protocol.close()
78
+ await self._protocol.wait_closed()
79
+ self._protocol = None
80
+ self._db_id = None
81
+
82
+ async def __aenter__(self) -> "DqliteConnection":
83
+ await self.connect()
84
+ return self
85
+
86
+ async def __aexit__(self, *args: Any) -> None:
87
+ await self.close()
88
+
89
+ def _ensure_connected(self) -> tuple[DqliteProtocol, int]:
90
+ """Ensure we're connected and return protocol and db_id."""
91
+ if self._protocol is None or self._db_id is None:
92
+ raise ConnectionError("Not connected")
93
+ return self._protocol, self._db_id
94
+
95
+ async def execute(self, sql: str, params: list[Any] | None = None) -> tuple[int, int]:
96
+ """Execute a SQL statement.
97
+
98
+ Returns (last_insert_id, rows_affected).
99
+ """
100
+ protocol, db_id = self._ensure_connected()
101
+ return await protocol.exec_sql(db_id, sql, params)
102
+
103
+ async def fetch(self, sql: str, params: list[Any] | None = None) -> list[dict[str, Any]]:
104
+ """Execute a query and return results as list of dicts."""
105
+ protocol, db_id = self._ensure_connected()
106
+ columns, rows = await protocol.query_sql(db_id, sql, params)
107
+ return [dict(zip(columns, row, strict=True)) for row in rows]
108
+
109
+ async def fetchall(self, sql: str, params: list[Any] | None = None) -> list[list[Any]]:
110
+ """Execute a query and return results as list of lists."""
111
+ protocol, db_id = self._ensure_connected()
112
+ _, rows = await protocol.query_sql(db_id, sql, params)
113
+ return rows
114
+
115
+ async def fetchone(self, sql: str, params: list[Any] | None = None) -> dict[str, Any] | None:
116
+ """Execute a query and return the first result."""
117
+ results = await self.fetch(sql, params)
118
+ return results[0] if results else None
119
+
120
+ async def fetchval(self, sql: str, params: list[Any] | None = None) -> Any:
121
+ """Execute a query and return the first column of the first row."""
122
+ protocol, db_id = self._ensure_connected()
123
+ _, rows = await protocol.query_sql(db_id, sql, params)
124
+ if rows and rows[0]:
125
+ return rows[0][0]
126
+ return None
127
+
128
+ @asynccontextmanager
129
+ async def transaction(self) -> AsyncIterator[None]:
130
+ """Context manager for transactions."""
131
+ if self._in_transaction:
132
+ # Nested transaction - just yield
133
+ yield
134
+ return
135
+
136
+ await self.execute("BEGIN")
137
+ self._in_transaction = True
138
+
139
+ try:
140
+ yield
141
+ await self.execute("COMMIT")
142
+ except Exception:
143
+ await self.execute("ROLLBACK")
144
+ raise
145
+ finally:
146
+ self._in_transaction = False
@@ -0,0 +1,37 @@
1
+ """Exceptions for dqlite client."""
2
+
3
+
4
+ class DqliteError(Exception):
5
+ """Base exception for dqlite client errors."""
6
+
7
+ pass
8
+
9
+
10
+ class ConnectionError(DqliteError):
11
+ """Error establishing or maintaining connection."""
12
+
13
+ pass
14
+
15
+
16
+ class ProtocolError(DqliteError):
17
+ """Protocol-level error."""
18
+
19
+ pass
20
+
21
+
22
+ class ClusterError(DqliteError):
23
+ """Cluster-related error (leader not found, etc)."""
24
+
25
+ pass
26
+
27
+
28
+ class OperationalError(DqliteError):
29
+ """Database operation error."""
30
+
31
+ code: int
32
+ message: str
33
+
34
+ def __init__(self, code: int, message: str) -> None:
35
+ self.code = code
36
+ self.message = message
37
+ super().__init__(f"[{code}] {message}")
@@ -0,0 +1,45 @@
1
+ """Node store interfaces for cluster discovery."""
2
+
3
+ from abc import ABC, abstractmethod
4
+ from dataclasses import dataclass
5
+
6
+
7
+ @dataclass
8
+ class NodeInfo:
9
+ """Information about a cluster node."""
10
+
11
+ node_id: int
12
+ address: str
13
+ role: int # 0=spare, 1=voter, 2=standby
14
+
15
+
16
+ class NodeStore(ABC):
17
+ """Abstract interface for storing cluster node information."""
18
+
19
+ @abstractmethod
20
+ async def get_nodes(self) -> list[NodeInfo]:
21
+ """Get list of known nodes."""
22
+ ...
23
+
24
+ @abstractmethod
25
+ async def set_nodes(self, nodes: list[NodeInfo]) -> None:
26
+ """Update list of known nodes."""
27
+ ...
28
+
29
+
30
+ class MemoryNodeStore(NodeStore):
31
+ """In-memory node store."""
32
+
33
+ def __init__(self, initial_addresses: list[str] | None = None) -> None:
34
+ self._nodes: list[NodeInfo] = []
35
+ if initial_addresses:
36
+ for i, addr in enumerate(initial_addresses):
37
+ self._nodes.append(NodeInfo(node_id=i, address=addr, role=1))
38
+
39
+ async def get_nodes(self) -> list[NodeInfo]:
40
+ """Get list of known nodes."""
41
+ return list(self._nodes)
42
+
43
+ async def set_nodes(self, nodes: list[NodeInfo]) -> None:
44
+ """Update list of known nodes."""
45
+ self._nodes = list(nodes)
dqliteclient/pool.py ADDED
@@ -0,0 +1,130 @@
1
+ """Connection pooling for dqlite."""
2
+
3
+ import asyncio
4
+ from collections.abc import AsyncIterator
5
+ from contextlib import asynccontextmanager
6
+ from typing import Any
7
+
8
+ from dqliteclient.cluster import ClusterClient
9
+ from dqliteclient.connection import DqliteConnection
10
+ from dqliteclient.exceptions import ConnectionError
11
+
12
+
13
+ class ConnectionPool:
14
+ """Connection pool with automatic leader detection."""
15
+
16
+ def __init__(
17
+ self,
18
+ addresses: list[str],
19
+ *,
20
+ database: str = "default",
21
+ min_size: int = 1,
22
+ max_size: int = 10,
23
+ timeout: float = 10.0,
24
+ ) -> None:
25
+ """Initialize connection pool.
26
+
27
+ Args:
28
+ addresses: List of node addresses
29
+ database: Database name
30
+ min_size: Minimum connections to maintain
31
+ max_size: Maximum connections allowed
32
+ timeout: Connection timeout
33
+ """
34
+ self._addresses = addresses
35
+ self._database = database
36
+ self._min_size = min_size
37
+ self._max_size = max_size
38
+ self._timeout = timeout
39
+
40
+ self._cluster = ClusterClient.from_addresses(addresses, timeout=timeout)
41
+ self._pool: asyncio.Queue[DqliteConnection] = asyncio.Queue(maxsize=max_size)
42
+ self._size = 0
43
+ self._lock = asyncio.Lock()
44
+ self._closed = False
45
+
46
+ async def initialize(self) -> None:
47
+ """Initialize the pool with minimum connections."""
48
+ for _ in range(self._min_size):
49
+ conn = await self._create_connection()
50
+ await self._pool.put(conn)
51
+
52
+ async def _create_connection(self) -> DqliteConnection:
53
+ """Create a new connection to the leader."""
54
+ conn = await self._cluster.connect(database=self._database)
55
+ self._size += 1
56
+ return conn
57
+
58
+ @asynccontextmanager
59
+ async def acquire(self) -> AsyncIterator[DqliteConnection]:
60
+ """Acquire a connection from the pool."""
61
+ if self._closed:
62
+ raise ConnectionError("Pool is closed")
63
+
64
+ conn: DqliteConnection | None = None
65
+
66
+ # Try to get from pool
67
+ try:
68
+ conn = self._pool.get_nowait()
69
+ except asyncio.QueueEmpty:
70
+ # Create new if under max
71
+ async with self._lock:
72
+ if self._size < self._max_size:
73
+ conn = await self._create_connection()
74
+
75
+ # Wait for one if at max
76
+ if conn is None:
77
+ conn = await self._pool.get()
78
+
79
+ try:
80
+ # Verify connection is still good
81
+ if not conn.is_connected:
82
+ await conn.connect()
83
+
84
+ yield conn
85
+ except Exception:
86
+ # On error, close connection and create new one
87
+ import contextlib
88
+
89
+ with contextlib.suppress(Exception):
90
+ await conn.close()
91
+ self._size -= 1
92
+ raise
93
+ else:
94
+ # Return to pool
95
+ try:
96
+ self._pool.put_nowait(conn)
97
+ except asyncio.QueueFull:
98
+ # Pool full, close connection
99
+ await conn.close()
100
+ self._size -= 1
101
+
102
+ async def execute(self, sql: str, params: list[Any] | None = None) -> tuple[int, int]:
103
+ """Execute a SQL statement using a pooled connection."""
104
+ async with self.acquire() as conn:
105
+ return await conn.execute(sql, params)
106
+
107
+ async def fetch(self, sql: str, params: list[Any] | None = None) -> list[dict[str, Any]]:
108
+ """Execute a query using a pooled connection."""
109
+ async with self.acquire() as conn:
110
+ return await conn.fetch(sql, params)
111
+
112
+ async def close(self) -> None:
113
+ """Close all connections in the pool."""
114
+ self._closed = True
115
+
116
+ while not self._pool.empty():
117
+ try:
118
+ conn = self._pool.get_nowait()
119
+ await conn.close()
120
+ except asyncio.QueueEmpty:
121
+ break
122
+
123
+ self._size = 0
124
+
125
+ async def __aenter__(self) -> "ConnectionPool":
126
+ await self.initialize()
127
+ return self
128
+
129
+ async def __aexit__(self, *args: Any) -> None:
130
+ await self.close()
@@ -0,0 +1,213 @@
1
+ """Low-level protocol handler for dqlite."""
2
+
3
+ import asyncio
4
+ from typing import Any
5
+
6
+ from dqliteclient.exceptions import ConnectionError, OperationalError, ProtocolError
7
+ from dqlitewire import MessageDecoder, MessageEncoder, ReadBuffer
8
+ from dqlitewire.messages import (
9
+ ClientRequest,
10
+ DbResponse,
11
+ ExecSqlRequest,
12
+ FailureResponse,
13
+ FinalizeRequest,
14
+ LeaderRequest,
15
+ LeaderResponse,
16
+ OpenRequest,
17
+ PrepareRequest,
18
+ QuerySqlRequest,
19
+ ResultResponse,
20
+ RowsResponse,
21
+ StmtResponse,
22
+ WelcomeResponse,
23
+ )
24
+ from dqlitewire.messages.base import Message
25
+
26
+
27
+ class DqliteProtocol:
28
+ """Low-level protocol handler for a single dqlite connection."""
29
+
30
+ def __init__(
31
+ self,
32
+ reader: asyncio.StreamReader,
33
+ writer: asyncio.StreamWriter,
34
+ ) -> None:
35
+ self._reader = reader
36
+ self._writer = writer
37
+ self._encoder = MessageEncoder()
38
+ self._decoder = MessageDecoder(is_request=False)
39
+ self._buffer = ReadBuffer()
40
+ self._client_id = 0
41
+ self._heartbeat_timeout = 0
42
+
43
+ async def handshake(self, client_id: int = 0) -> int:
44
+ """Perform protocol handshake.
45
+
46
+ Returns the heartbeat timeout from server.
47
+ """
48
+ # Send protocol version
49
+ self._writer.write(self._encoder.encode_handshake())
50
+ await self._writer.drain()
51
+
52
+ # Send client registration
53
+ request = ClientRequest(client_id=client_id)
54
+ self._writer.write(request.encode())
55
+ await self._writer.drain()
56
+
57
+ # Read welcome response
58
+ response = await self._read_response()
59
+
60
+ if isinstance(response, FailureResponse):
61
+ raise ProtocolError(f"Handshake failed: {response.message}")
62
+
63
+ if not isinstance(response, WelcomeResponse):
64
+ raise ProtocolError(f"Expected WelcomeResponse, got {type(response).__name__}")
65
+
66
+ self._client_id = client_id
67
+ self._heartbeat_timeout = response.heartbeat_timeout
68
+ return response.heartbeat_timeout
69
+
70
+ async def get_leader(self) -> tuple[int, str]:
71
+ """Request leader information.
72
+
73
+ Returns (node_id, address).
74
+ """
75
+ request = LeaderRequest()
76
+ self._writer.write(request.encode())
77
+ await self._writer.drain()
78
+
79
+ response = await self._read_response()
80
+
81
+ if isinstance(response, FailureResponse):
82
+ raise OperationalError(response.code, response.message)
83
+
84
+ if not isinstance(response, LeaderResponse):
85
+ raise ProtocolError(f"Expected LeaderResponse, got {type(response).__name__}")
86
+
87
+ return response.node_id, response.address
88
+
89
+ async def open_database(self, name: str, flags: int = 0, vfs: str = "") -> int:
90
+ """Open a database.
91
+
92
+ Returns the database ID.
93
+ """
94
+ request = OpenRequest(name=name, flags=flags, vfs=vfs)
95
+ self._writer.write(request.encode())
96
+ await self._writer.drain()
97
+
98
+ response = await self._read_response()
99
+
100
+ if isinstance(response, FailureResponse):
101
+ raise OperationalError(response.code, response.message)
102
+
103
+ if not isinstance(response, DbResponse):
104
+ raise ProtocolError(f"Expected DbResponse, got {type(response).__name__}")
105
+
106
+ return response.db_id
107
+
108
+ async def prepare(self, db_id: int, sql: str) -> tuple[int, int]:
109
+ """Prepare a SQL statement.
110
+
111
+ Returns (stmt_id, num_params).
112
+ """
113
+ request = PrepareRequest(db_id=db_id, sql=sql)
114
+ self._writer.write(request.encode())
115
+ await self._writer.drain()
116
+
117
+ response = await self._read_response()
118
+
119
+ if isinstance(response, FailureResponse):
120
+ raise OperationalError(response.code, response.message)
121
+
122
+ if not isinstance(response, StmtResponse):
123
+ raise ProtocolError(f"Expected StmtResponse, got {type(response).__name__}")
124
+
125
+ return response.stmt_id, response.num_params
126
+
127
+ async def finalize(self, db_id: int, stmt_id: int) -> None:
128
+ """Finalize (close) a prepared statement."""
129
+ request = FinalizeRequest(db_id=db_id, stmt_id=stmt_id)
130
+ self._writer.write(request.encode())
131
+ await self._writer.drain()
132
+
133
+ response = await self._read_response()
134
+
135
+ if isinstance(response, FailureResponse):
136
+ raise OperationalError(response.code, response.message)
137
+
138
+ async def exec_sql(
139
+ self, db_id: int, sql: str, params: list[Any] | None = None
140
+ ) -> tuple[int, int]:
141
+ """Execute SQL directly.
142
+
143
+ Returns (last_insert_id, rows_affected).
144
+ """
145
+ request = ExecSqlRequest(db_id=db_id, sql=sql, params=params or [])
146
+ self._writer.write(request.encode())
147
+ await self._writer.drain()
148
+
149
+ response = await self._read_response()
150
+
151
+ if isinstance(response, FailureResponse):
152
+ raise OperationalError(response.code, response.message)
153
+
154
+ if not isinstance(response, ResultResponse):
155
+ raise ProtocolError(f"Expected ResultResponse, got {type(response).__name__}")
156
+
157
+ return response.last_insert_id, response.rows_affected
158
+
159
+ async def query_sql(
160
+ self, db_id: int, sql: str, params: list[Any] | None = None
161
+ ) -> tuple[list[str], list[list[Any]]]:
162
+ """Execute a query directly.
163
+
164
+ Returns (column_names, rows).
165
+ """
166
+ request = QuerySqlRequest(db_id=db_id, sql=sql, params=params or [])
167
+ self._writer.write(request.encode())
168
+ await self._writer.drain()
169
+
170
+ response = await self._read_response()
171
+
172
+ if isinstance(response, FailureResponse):
173
+ raise OperationalError(response.code, response.message)
174
+
175
+ if not isinstance(response, RowsResponse):
176
+ raise ProtocolError(f"Expected RowsResponse, got {type(response).__name__}")
177
+
178
+ # Store column names from first response
179
+ column_names = response.column_names
180
+
181
+ # Handle multi-part responses
182
+ all_rows = list(response.rows)
183
+ while response.has_more:
184
+ next_response = await self._read_response()
185
+ if isinstance(next_response, RowsResponse):
186
+ all_rows.extend(next_response.rows)
187
+ response = next_response
188
+ else:
189
+ break
190
+
191
+ return column_names, all_rows
192
+
193
+ async def _read_response(self) -> Message:
194
+ """Read and decode the next response message."""
195
+ while not self._decoder.has_message():
196
+ data = await self._reader.read(4096)
197
+ if not data:
198
+ raise ConnectionError("Connection closed by server")
199
+ self._decoder.feed(data)
200
+
201
+ message = self._decoder.decode()
202
+ if message is None:
203
+ raise ProtocolError("Failed to decode message")
204
+
205
+ return message
206
+
207
+ def close(self) -> None:
208
+ """Close the connection."""
209
+ self._writer.close()
210
+
211
+ async def wait_closed(self) -> None:
212
+ """Wait for the connection to close."""
213
+ await self._writer.wait_closed()
dqliteclient/py.typed ADDED
File without changes
dqliteclient/retry.py ADDED
@@ -0,0 +1,51 @@
1
+ """Retry utilities with exponential backoff."""
2
+
3
+ import asyncio
4
+ import random
5
+ from collections.abc import Awaitable, Callable
6
+
7
+
8
+ async def retry_with_backoff[T](
9
+ func: Callable[[], Awaitable[T]],
10
+ max_attempts: int = 5,
11
+ base_delay: float = 0.1,
12
+ max_delay: float = 10.0,
13
+ jitter: float = 0.1,
14
+ ) -> T:
15
+ """Retry an async function with exponential backoff.
16
+
17
+ Args:
18
+ func: Async function to retry
19
+ max_attempts: Maximum number of attempts
20
+ base_delay: Initial delay between retries in seconds
21
+ max_delay: Maximum delay between retries
22
+ jitter: Random jitter factor (0-1)
23
+
24
+ Returns:
25
+ Result of the function
26
+
27
+ Raises:
28
+ The last exception if all attempts fail
29
+ """
30
+ last_error: Exception | None = None
31
+
32
+ for attempt in range(max_attempts):
33
+ try:
34
+ return await func()
35
+ except Exception as e:
36
+ last_error = e
37
+
38
+ if attempt == max_attempts - 1:
39
+ break
40
+
41
+ # Calculate delay with exponential backoff
42
+ delay = min(base_delay * (2**attempt), max_delay)
43
+
44
+ # Add jitter
45
+ if jitter > 0:
46
+ delay = delay * (1 + random.uniform(-jitter, jitter))
47
+
48
+ await asyncio.sleep(delay)
49
+
50
+ assert last_error is not None
51
+ raise last_error