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
4 changes: 4 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0

## [Unreleased]

### Changed

- `ProtocolVersion` enum renamed to `StartupRequestCode`. The `version_code` parameter on `StartupMessageRegistry.register()` and `.lookup()` is now `request_code`.

## [0.1.0] - 2026-03-11

### Added
Expand Down
4 changes: 2 additions & 2 deletions docs/reference/constants.md
Original file line number Diff line number Diff line change
Expand Up @@ -33,9 +33,9 @@ Protocol-level enums and constants.
| `FRONTEND` | `"frontend"` | Message sent by client |
| `BACKEND` | `"backend"` | Message sent by server |

## `ProtocolVersion`
## `StartupRequestCode`

`IntEnum` of PostgreSQL protocol version codes. Used in `StartupMessage` and special request messages.
`IntEnum` of 32-bit codes sent in the startup packet request code field. Used in `StartupMessage` and special request messages.

| Member | Value | Description |
|--------|-------|-------------|
Expand Down
2 changes: 1 addition & 1 deletion docs/reference/messages/startup.md
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ Connection initialization. Sent as the first message from a client.
| Field | Type | Description |
|-------|------|-------------|
| `params` | `dict[str, str]` | Key-value parameters (`user`, `database`, etc.) |
| `protocol_version` | `int` | Protocol version code (default: `ProtocolVersion.V3_0`) |
| `protocol_version` | `int` | Startup request code (default: `StartupRequestCode.V3_0`) |

```python
from pygwire.messages import StartupMessage
Expand Down
4 changes: 2 additions & 2 deletions src/pygwire/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@
Connection,
FrontendConnection,
)
from pygwire.constants import ConnectionPhase, ProtocolVersion, TransactionStatus
from pygwire.constants import ConnectionPhase, StartupRequestCode, TransactionStatus
from pygwire.exceptions import (
DecodingError,
FramingError,
Expand All @@ -35,7 +35,7 @@
"FrontendMessageDecoder",
"FrontendStateMachine",
"ProtocolError",
"ProtocolVersion",
"StartupRequestCode",
"PygwireError",
"StateMachineError",
"TransactionStatus",
Expand Down
11 changes: 8 additions & 3 deletions src/pygwire/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
__all__ = [
"ConnectionPhase",
"MessageDirection",
"ProtocolVersion",
"StartupRequestCode",
"TransactionStatus",
]
Comment thread
DHUKK marked this conversation as resolved.

Expand Down Expand Up @@ -80,8 +80,13 @@ class ConnectionPhase(Enum):
FAILED = auto()


class ProtocolVersion(IntEnum):
"""PostgreSQL Protocol Versions."""
class StartupRequestCode(IntEnum):
"""32-bit codes sent in the startup packet version field.

The first 4 bytes of every startup packet are read as a request code.
V3_0 and V3_2 are actual protocol versions; SSL_REQUEST, GSSENC_REQUEST,
and CANCEL_REQUEST are magic numbers that share the same wire position.
"""

V3_0 = 0x00030000 # Standard for PG 14-17
V3_2 = 0x00030002 # New for PG 18+ (Variable length cancel keys)
Expand Down
14 changes: 7 additions & 7 deletions src/pygwire/framing.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,12 +97,12 @@ class StartupFraming(FramingStrategy):

Wire format:
Bytes 0-3: Int32 length (including these 4 bytes)
Bytes 4-7: Int32 version_code (part of payload)
Bytes 4-7: Int32 request_code (part of payload)
Bytes 8+: Remaining payload

Example:
StartupMessage: length=52, version_code=0x00030000, params...
SSLRequest: length=8, version_code=80877103
StartupMessage: length=52, request_code=0x00030000, params...
SSLRequest: length=8, request_code=80877103
"""

def try_parse(
Expand All @@ -129,13 +129,13 @@ def try_parse(
payload_end = pos + length
payload = buf[payload_start:payload_end]
if len(payload) < 4:
raise FramingError("Startup message payload too short for version code")
raise FramingError("Startup message payload too short for request code")

(version_code,) = _LENGTH_STRUCT.unpack_from(payload)
(request_code,) = _LENGTH_STRUCT.unpack_from(payload)

msg_cls = STARTUP_REGISTRY.lookup(version_code)
msg_cls = STARTUP_REGISTRY.lookup(request_code)
if msg_cls is None:
raise FramingError(f"Unknown startup message version code: {version_code:#010x}")
raise FramingError(f"Unknown startup message request code: {request_code:#010x}")
try:
msg = msg_cls.decode(payload)
except struct.error as e:
Expand Down
26 changes: 13 additions & 13 deletions src/pygwire/messages/_registry.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,11 @@
"""Message registry system for PostgreSQL wire protocol.

This module provides the registry infrastructure that maps message identifiers
and version codes to message classes. It supports three types of registries
and request codes to message classes. It supports three types of registries
for different framing modes:

- StandardMessageRegistry: Standard framed messages (Byte1 + Int32 + payload)
- StartupMessageRegistry: Startup messages (Int32 + payload, version code discriminator)
- StartupMessageRegistry: Startup messages (Int32 + payload, request code discriminator)
- NegotiationMessageRegistry: SSL/GSS negotiation (single byte messages)
"""

Expand Down Expand Up @@ -107,9 +107,9 @@ def lookup(


class StartupMessageRegistry:
"""Registry for startup messages (Int32 + payload, version code discriminator).
"""Registry for startup messages (Int32 + payload, request code discriminator).

Startup messages have no identifier byte. Instead, they use a 4-byte version
Startup messages have no identifier byte. Instead, they use a 4-byte request
code at the start of the payload to distinguish message types:
- 0x00030000: StartupMessage
- 80877103: SSLRequest
Expand All @@ -120,38 +120,38 @@ class StartupMessageRegistry:
"""

def __init__(self) -> None:
# Key: version_code → message class
# Key: request_code → message class
self._registry: dict[int, type[PGMessage]] = {}

def register(self, version_code: int) -> Callable[[type[PGMessage]], type[PGMessage]]:
def register(self, request_code: int) -> Callable[[type[PGMessage]], type[PGMessage]]:
"""Decorator to register a startup message class.

Args:
version_code: 32-bit version/request code
request_code: 32-bit version/request code

Example::

@STARTUP_REGISTRY.register(version_code=0x00030000)
@STARTUP_REGISTRY.register(request_code=0x00030000)
class StartupMessage(SpecialMessage):
...
"""

def decorator(cls: type[PGMessage]) -> type[PGMessage]:
self._registry[version_code] = cls
self._registry[request_code] = cls
return cls

return decorator

def lookup(self, version_code: int) -> type[PGMessage] | None:
"""Find message class by version code.
def lookup(self, request_code: int) -> type[PGMessage] | None:
"""Find message class by request code.

Args:
version_code: 32-bit version/request code
request_code: 32-bit version/request code

Returns:
Message class or None if not found
"""
return self._registry.get(version_code)
return self._registry.get(request_code)


class NegotiationMessageRegistry:
Expand Down
22 changes: 11 additions & 11 deletions src/pygwire/messages/_startup.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,31 +6,31 @@
from dataclasses import dataclass, field
from typing import Self

from pygwire.constants import ProtocolVersion
from pygwire.constants import StartupRequestCode

from ._base import SpecialMessage, _read_cstring
from ._registry import STARTUP_REGISTRY

_INT32 = struct.Struct("!I")


@STARTUP_REGISTRY.register(version_code=ProtocolVersion.V3_0)
@STARTUP_REGISTRY.register(version_code=ProtocolVersion.V3_2)
@STARTUP_REGISTRY.register(request_code=StartupRequestCode.V3_0)
@STARTUP_REGISTRY.register(request_code=StartupRequestCode.V3_2)
@dataclass(slots=True)
class StartupMessage(SpecialMessage):
"""StartupMessage — initial connection packet (Protocol 3.0 & 3.2).

Contains key-value parameters (user, database, options, etc.)
terminated by a final null byte. The ``encode()`` method returns
the full payload including the Int32 version code.
the full payload including the Int32 request code.

Note: The message format is identical in v3.0 and v3.2. Protocol version
3.2 (PostgreSQL 18+) only differs in CancelRequest and BackendKeyData
messages which support variable-length secret keys.
"""

params: dict[str, str] = field(default_factory=dict)
protocol_version: int = ProtocolVersion.V3_0
protocol_version: int = StartupRequestCode.V3_0

def encode(self) -> bytes:
buf = bytearray(_INT32.pack(self.protocol_version))
Expand All @@ -54,7 +54,7 @@ def decode(cls, payload: memoryview) -> Self:
return cls(params=params, protocol_version=protocol_version)


@STARTUP_REGISTRY.register(version_code=ProtocolVersion.SSL_REQUEST)
@STARTUP_REGISTRY.register(request_code=StartupRequestCode.SSL_REQUEST)
@dataclass(slots=True)
class SSLRequest(SpecialMessage):
"""SSLRequest — asks if the server supports SSL.
Expand All @@ -63,14 +63,14 @@ class SSLRequest(SpecialMessage):
"""

def encode(self) -> bytes:
return _INT32.pack(ProtocolVersion.SSL_REQUEST)
return _INT32.pack(StartupRequestCode.SSL_REQUEST)

@classmethod
def decode(cls, payload: memoryview) -> Self:
return cls()


@STARTUP_REGISTRY.register(version_code=ProtocolVersion.GSSENC_REQUEST)
@STARTUP_REGISTRY.register(request_code=StartupRequestCode.GSSENC_REQUEST)
@dataclass(slots=True)
class GSSEncRequest(SpecialMessage):
"""GSSEncRequest — asks if the server supports GSS encryption.
Expand All @@ -79,14 +79,14 @@ class GSSEncRequest(SpecialMessage):
"""

def encode(self) -> bytes:
return _INT32.pack(ProtocolVersion.GSSENC_REQUEST)
return _INT32.pack(StartupRequestCode.GSSENC_REQUEST)

@classmethod
def decode(cls, payload: memoryview) -> Self:
return cls()


@STARTUP_REGISTRY.register(version_code=ProtocolVersion.CANCEL_REQUEST)
@STARTUP_REGISTRY.register(request_code=StartupRequestCode.CANCEL_REQUEST)
@dataclass(slots=True)
class CancelRequest(SpecialMessage):
"""CancelRequest — asks the server to cancel a running query.
Expand All @@ -101,7 +101,7 @@ class CancelRequest(SpecialMessage):

def encode(self) -> bytes:
return (
_INT32.pack(ProtocolVersion.CANCEL_REQUEST)
_INT32.pack(StartupRequestCode.CANCEL_REQUEST)
+ _INT32.pack(self.process_id)
+ self.secret_key
)
Expand Down
8 changes: 4 additions & 4 deletions tests/integration/test_malformed_payloads.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
import pytest

from pygwire.codec import BackendMessageDecoder, FrontendMessageDecoder
from pygwire.constants import ProtocolVersion
from pygwire.constants import StartupRequestCode
from pygwire.exceptions import DecodingError, FramingError
from pygwire.messages import (
AuthenticationSASL,
Expand Down Expand Up @@ -165,7 +165,7 @@ def test_error_response_truncated_fields(self):
next(decoder)

def test_cancel_request_truncated(self):
"""Test CancelRequest with valid version code but truncated payload.
"""Test CancelRequest with valid request code but truncated payload.

CancelRequest uses startup framing (no identifier byte). This exercises
the StartupFraming struct.error catch path when decode fails on a
Expand All @@ -176,8 +176,8 @@ def test_cancel_request_truncated(self):

# CancelRequest wire format: Int32(length) + Int32(cancel_code) + Int32(pid) + secret_key
# Craft a message with valid cancel code but truncated before process_id
cancel_code = int(ProtocolVersion.CANCEL_REQUEST)
payload = struct.pack("!I", cancel_code) # version code only, no pid/key
cancel_code = int(StartupRequestCode.CANCEL_REQUEST)
payload = struct.pack("!I", cancel_code) # request code only, no pid/key
length = 4 + len(payload) # length includes itself
wire = struct.pack("!I", length) + payload

Expand Down
14 changes: 7 additions & 7 deletions tests/unit/test_framing.py
Original file line number Diff line number Diff line change
Expand Up @@ -102,24 +102,24 @@ def test_message_exceeds_max_size(self):
memoryview(wire), 0, ConnectionPhase.STARTUP, MessageDirection.FRONTEND
)

def test_payload_too_short_for_version_code(self):
def test_payload_too_short_for_request_code(self):
"""Test that payload shorter than 4 bytes raises FramingError."""
# Length = 5 (header) + 2 (payload) = 7, but payload needs 4 bytes for version
wire = struct.pack("!I", 6) + b"ab"

framing = StartupFraming()
with pytest.raises(FramingError, match="payload too short for version code"):
with pytest.raises(FramingError, match="payload too short for request code"):
framing.try_parse(
memoryview(wire), 0, ConnectionPhase.STARTUP, MessageDirection.FRONTEND
)

def test_unknown_version_code_raises_error(self):
"""Test that unknown version code raises FramingError."""
# Create message with invalid version code
wire = struct.pack("!II", 8, 0xDEADBEEF) # Invalid version code
def test_unknown_request_code_raises_error(self):
"""Test that unknown request code raises FramingError."""
# Create message with invalid request code
wire = struct.pack("!II", 8, 0xDEADBEEF) # Invalid request code

framing = StartupFraming()
with pytest.raises(FramingError, match="Unknown startup message version code"):
with pytest.raises(FramingError, match="Unknown startup message request code"):
framing.try_parse(
memoryview(wire), 0, ConnectionPhase.STARTUP, MessageDirection.FRONTEND
)
Expand Down
Loading
Loading