Skip to content
Merged
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
15 changes: 1 addition & 14 deletions memcache/async_connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
import anyio
from anyio.streams.buffered import BufferedByteReceiveStream

from .errors import MemcacheError
from .errors import MemcacheError, PipelineError as PipelineError
from .meta_command import MetaCommand, MetaResult


Expand Down Expand Up @@ -120,19 +120,6 @@ async def _receive_meta_result(self) -> MetaResult:
return result


class PipelineError(Exception):
def __init__(
self,
written: int,
responses: List[MetaResult],
cause: BaseException,
) -> None:
super().__init__(str(cause))
self.written = written
self.responses = responses
self.cause = cause


class AsyncPool:
def __init__(
self,
Expand Down
9 changes: 9 additions & 0 deletions memcache/async_memcache.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,15 @@ def __init__(
password=password,
)

async def __aenter__(self) -> "AsyncMemcache":
return self

async def __aexit__(self, *exc: Any) -> None:
await self.close()

async def close(self) -> None:
await self._meta.close()

@asynccontextmanager
async def _get_connection(
self, key: Union[str, bytes]
Expand Down
96 changes: 73 additions & 23 deletions memcache/connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,9 @@
import socket
import threading
from contextlib import contextmanager
from typing import Callable, Iterator, Optional, Tuple
from typing import Callable, Iterator, List, Optional, Tuple

from .errors import MemcacheError
from .errors import MemcacheError, PipelineError
from .meta_command import MetaCommand, MetaResult


Expand All @@ -20,17 +20,26 @@ def __init__(
*,
username: Optional[str] = None,
password: Optional[str] = None,
timeout: Optional[float] = None,
):
self._addr = addr
self._username = username
self._password = password
self._connect()
self._connect(timeout)

def _connect(self) -> None:
def _connect(self, timeout: Optional[float]) -> None:
self.socket = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
self.socket.connect(self._addr)
self.stream = self.socket.makefile(mode="rwb")
self._auth()
self.socket.settimeout(timeout)
try:
self.socket.connect(self._addr)
self.stream = self.socket.makefile(mode="rb")
self._auth()
except BaseException:
self.socket.close()
raise

def _set_timeout(self, timeout: Optional[float]) -> None:
self.socket.settimeout(timeout)

def _auth(self) -> None:
if self._username is None or self._password is None:
Expand All @@ -39,42 +48,46 @@ def _auth(self) -> None:
self._username.encode("utf-8"),
self._password.encode("utf-8"),
)
self.stream.write(b"set auth x 0 %d\r\n" % len(auth_data))
self.stream.write(auth_data)
self.stream.write(b"\r\n")
self.stream.flush()
self.socket.sendall(
b"set auth x 0 %d\r\n" % len(auth_data) + auth_data + b"\r\n"
)
response = self.stream.readline()
if response != b"STORED\r\n":
raise MemcacheError(response.rstrip(NEWLINE))

def close(self) -> None:
self.stream.close()
self.socket.close()
try:
self.stream.close()
finally:
self.socket.close()

def flush_all(self, delay: int = 0) -> None:
def flush_all(self, delay: int = 0, timeout: Optional[float] = None) -> None:
self._set_timeout(timeout)
if delay > 0:
self.stream.write(b"flush_all %d\r\n" % delay)
self.socket.sendall(b"flush_all %d\r\n" % delay)
else:
self.stream.write(b"flush_all\r\n")
self.stream.flush()
self.socket.sendall(b"flush_all\r\n")
response = self.stream.readline()
if response != b"OK\r\n":
raise MemcacheError(response.rstrip(NEWLINE))

def execute_meta_command(self, command: MetaCommand) -> MetaResult:
def execute_meta_command(
self, command: MetaCommand, timeout: Optional[float] = None
) -> MetaResult:
# Never reconnect and replay here. Once a write has started, a lost
# response makes the outcome ambiguous (especially for ms/ma).
self._set_timeout(timeout)
return self._execute_meta_command(command)

def _execute_meta_command(self, command: MetaCommand) -> MetaResult:
self.stream.write(command.dump_header())
if command.value:
self.stream.write(command.value + b"\r\n")
self.stream.flush()
self.socket.sendall(command.dump())
return self._receive_meta_result()

def _receive_meta_result(self) -> MetaResult:
result = MetaResult.load_header(self.stream.readline())
line = self.stream.readline()
if not line:
raise MemcacheError("connection closed while reading response")
result = MetaResult.load_header(line)

if result.rc == b"VA":
if result.datalen is None:
Expand All @@ -84,6 +97,34 @@ def _receive_meta_result(self) -> MetaResult:

return result

def execute_pipeline(
self, commands: List[MetaCommand], timeout: Optional[float] = None
) -> List[MetaResult]:
"""Write a quiet pipeline and read through its ``mn`` barrier."""
self._set_timeout(timeout)
written = 0
responses: List[MetaResult] = []
try:
for command in commands:
written += 1
self.socket.sendall(command.dump())
self.socket.sendall(b"mn\r\n")
while True:
line = self.stream.readline()
if not line:
raise MemcacheError("connection closed while reading pipeline")
if line == b"MN\r\n":
return responses
result = MetaResult.load_header(line)
if result.rc == b"VA":
if result.datalen is None:
raise MemcacheError("invalid response: missing datalen")
result.value = self.stream.read(result.datalen)
self.stream.read(2)
responses.append(result)
except BaseException as exc:
raise PipelineError(written, responses, exc)


class Pool:
def __init__(
Expand Down Expand Up @@ -126,3 +167,12 @@ def get(self) -> Iterator[Connection]:
raise
else:
self._connections.put(connection)

def close(self) -> None:
while True:
try:
connection = self._connections.get_nowait()
except queue.Empty:
break
connection.close()
self._size = 0
22 changes: 21 additions & 1 deletion memcache/errors.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,9 @@
from typing import Any, Optional
from __future__ import annotations

from typing import TYPE_CHECKING, Any, List, Optional

if TYPE_CHECKING:
from .meta_command import MetaResult


class MemcacheError(Exception):
Expand All @@ -19,3 +24,18 @@ def __init__(self, result: Optional[Any] = None) -> None:

class ProtocolError(MemcacheError):
"""The server returned a malformed or unsupported protocol response."""


class PipelineError(MemcacheError):
"""A pipeline failed after a possibly partial write or response sequence."""

def __init__(
self,
written: int,
responses: List[MetaResult],
cause: BaseException,
) -> None:
super().__init__(str(cause))
self.written = written
self.responses = responses
self.cause = cause
Loading
Loading