diff --git a/changelogs/fragments/ssh-agent-hardening.yml b/changelogs/fragments/ssh-agent-hardening.yml new file mode 100644 index 00000000000..8e0daab293b --- /dev/null +++ b/changelogs/fragments/ssh-agent-hardening.yml @@ -0,0 +1,4 @@ +bugfixes: + - ssh-agent - fix partial socket reads + - ssh-agent - fix wire format serialization for zero-length values + - ssh-agent - cap agent response size to match OpenSSH limits diff --git a/lib/ansible/_internal/_ssh/_ssh_agent.py b/lib/ansible/_internal/_ssh/_ssh_agent.py index 345284d1b69..cc610f04533 100644 --- a/lib/ansible/_internal/_ssh/_ssh_agent.py +++ b/lib/ansible/_internal/_ssh/_ssh_agent.py @@ -4,7 +4,6 @@ from __future__ import annotations import binascii -import copy import dataclasses import enum import functools @@ -65,6 +64,7 @@ if t.TYPE_CHECKING: _SSH_AGENT_CLIENT_SOCKET_TIMEOUT = 10 +_SSH_AGENT_MAX_RESPONSE_BYTES = 256 * 1024 # not part of the RFC, just a safety measure class ProtocolMsgNumbers(enum.IntEnum): @@ -159,8 +159,8 @@ class mpint(int, VariableSized): def to_blob(self) -> bytes: if self < 0: raise ValueError("negative mpint not allowed") - if not self: - return b"" + if self == 0: + return uint32(self).to_blob() nbytes = (self.bit_length() + 8) // 8 ret = bytearray(self.to_bytes(length=nbytes, byteorder='big')) ret[:0] = uint32(len(ret)).to_blob() @@ -180,10 +180,7 @@ class constraints(bytes): class binary_string(bytes, VariableSized): def to_blob(self) -> bytes: - if length := len(self): - return uint32(length).to_blob() + self - else: - return b"" + return uint32(len(self)).to_blob() + self @classmethod def from_blob(cls, blob: memoryview | bytes) -> t.Self: @@ -193,17 +190,14 @@ class binary_string(bytes, VariableSized): class unicode_string(str, VariableSized): def to_blob(self) -> bytes: val = self.encode('utf-8') - if length := len(val): - return uint32(length).to_blob() + val - else: - return b"" + return uint32(len(val)).to_blob() + val @classmethod def from_blob(cls, blob: memoryview | bytes) -> t.Self: return cls(bytes(blob).decode('utf-8')) -class KeyAlgo(str, VariableSized, enum.Enum): +class KeyAlgo(VariableSized, enum.StrEnum): RSA = "ssh-rsa" DSA = "ssh-dss" ECDSA256 = "ecdsa-sha2-nistp256" @@ -338,7 +332,7 @@ class RSAPrivateKeyMsg(PrivateKeyMsg): iqmp: mpint p: mpint q: mpint - comments: unicode_string = dataclasses.field(default=unicode_string(''), compare=False) + comment: unicode_string = dataclasses.field(default=unicode_string(''), compare=False) constraints: constraints = dataclasses.field(default=constraints(b'')) @@ -350,7 +344,7 @@ class DSAPrivateKeyMsg(PrivateKeyMsg): g: mpint y: mpint x: mpint - comments: unicode_string = dataclasses.field(default=unicode_string(''), compare=False) + comment: unicode_string = dataclasses.field(default=unicode_string(''), compare=False) constraints: constraints = dataclasses.field(default=constraints(b'')) @@ -360,7 +354,7 @@ class EcdsaPrivateKeyMsg(PrivateKeyMsg): ecdsa_curve_name: unicode_string Q: binary_string d: mpint - comments: unicode_string = dataclasses.field(default=unicode_string(''), compare=False) + comment: unicode_string = dataclasses.field(default=unicode_string(''), compare=False) constraints: constraints = dataclasses.field(default=constraints(b'')) @@ -369,7 +363,7 @@ class Ed25519PrivateKeyMsg(PrivateKeyMsg): type: KeyAlgo enc_a: binary_string k_env_a: binary_string - comments: unicode_string = dataclasses.field(default=unicode_string(''), compare=False) + comment: unicode_string = dataclasses.field(default=unicode_string(''), compare=False) constraints: constraints = dataclasses.field(default=constraints(b'')) @@ -453,12 +447,7 @@ class PublicKeyMsg(Msg): @functools.cached_property def fingerprint(self) -> str: - digest = hashlib.sha256() - msg = copy.copy(self) - msg.comments = unicode_string('') - k = msg.to_blob() - digest.update(k) - return binascii.b2a_base64(digest.digest(), newline=False).rstrip(b'=').decode('utf-8') + return binascii.b2a_base64(hashlib.sha256(self.to_blob()).digest(), newline=False).rstrip(b'=').decode('utf-8') @dataclasses.dataclass(order=True, slots=True) @@ -466,7 +455,6 @@ class RSAPublicKeyMsg(PublicKeyMsg): type: KeyAlgo e: mpint n: mpint - comments: unicode_string = dataclasses.field(default=unicode_string(''), compare=False) @dataclasses.dataclass(order=True, slots=True) @@ -476,7 +464,6 @@ class DSAPublicKeyMsg(PublicKeyMsg): q: mpint g: mpint y: mpint - comments: unicode_string = dataclasses.field(default=unicode_string(''), compare=False) @dataclasses.dataclass(order=True, slots=True) @@ -484,61 +471,18 @@ class EcdsaPublicKeyMsg(PublicKeyMsg): type: KeyAlgo ecdsa_curve_name: unicode_string Q: binary_string - comments: unicode_string = dataclasses.field(default=unicode_string(''), compare=False) @dataclasses.dataclass(order=True, slots=True) class Ed25519PublicKeyMsg(PublicKeyMsg): type: KeyAlgo enc_a: binary_string - comments: unicode_string = dataclasses.field(default=unicode_string(''), compare=False) -@dataclasses.dataclass(order=True, slots=True) -class KeyList(Msg): - nkeys: uint32 - keys: PublicKeyMsgList - - def __post_init__(self) -> None: - if self.nkeys != len(self.keys): - raise SshAgentFailure("agent: invalid number of keys received for identities list") - - -@dataclasses.dataclass(order=True, slots=True) -class PublicKeyMsgList(Msg): - keys: list[PublicKeyMsg] - - def __iter__(self) -> t.Iterator[PublicKeyMsg]: - yield from self.keys - - def __len__(self) -> int: - return len(self.keys) - - @classmethod - def from_blob(cls, blob: memoryview | bytes) -> t.Self: ... - - @classmethod - def consume_from_blob(cls, blob: memoryview | bytes) -> tuple[t.Self, memoryview | bytes]: - args: list[PublicKeyMsg] = [] - while blob: - prev_blob = blob - key_blob, key_blob_length, comment_blob = cls._consume_field(blob) - - peek_key_algo, _length, _blob = cls._consume_field(key_blob) - pub_key_msg_cls = PublicKeyMsg.get_dataclass(KeyAlgo(bytes(peek_key_algo).decode('utf-8'))) - - _fv, comment_blob_length, blob = cls._consume_field(comment_blob) - key_plus_comment = prev_blob[4 : (4 + key_blob_length) + (4 + comment_blob_length)] - - args.append(pub_key_msg_cls.from_blob(key_plus_comment)) - return cls(args), b"" - - @staticmethod - def _consume_field(blob: memoryview | bytes) -> tuple[memoryview | bytes, uint32, memoryview | bytes]: - length = uint32.from_blob(blob[:4]) - blob = blob[4:] - data, rest = _split_blob(blob, length) - return data, length, rest +@dataclasses.dataclass(order=True, frozen=True, slots=True) +class Identity: + key: PublicKeyMsg + comment: unicode_string class SshAgentClient: @@ -561,11 +505,23 @@ class SshAgentClient: ) -> None: self.close() + def _read_all(self, bytes_to_read: int) -> bytes: + data_read = bytearray() + while bytes_to_read: + data = self._sock.recv(bytes_to_read) + if not data: + raise ConnectionError("agent: connection closed") + bytes_to_read -= len(data) + data_read.extend(data) + return bytes(data_read) + def send(self, msg: bytes) -> bytes: length = uint32(len(msg)).to_blob() self._sock.sendall(length + msg) - bufsize = uint32.from_blob(self._sock.recv(4)) - resp = self._sock.recv(bufsize) + bufsize = uint32.from_blob(self._read_all(4)) + if bufsize > _SSH_AGENT_MAX_RESPONSE_BYTES: + raise SshAgentFailure("agent: response too large") + resp = self._read_all(bufsize) if resp[0] == ProtocolMsgNumbers.SSH_AGENT_FAILURE: raise SshAgentFailure('agent: failure') return resp @@ -580,12 +536,12 @@ class SshAgentClient: def add( self, private_key: CryptoPrivateKey, - comments: str | None = None, + comment: str | None = None, lifetime: int | None = None, confirm: bool | None = None, ) -> None: key_msg = PrivateKeyMsg.from_private_key(private_key) - key_msg.comments = unicode_string(comments or '') + key_msg.comment = unicode_string(comment or '') if lifetime: key_msg.constraints += constraints([ProtocolMsgNumbers.SSH_AGENT_CONSTRAIN_LIFETIME]).to_blob() + uint32(lifetime).to_blob() if confirm: @@ -598,16 +554,40 @@ class SshAgentClient: msg += key_msg.to_blob() self.send(msg) - def list(self) -> KeyList: + def list(self) -> list[Identity]: req = ProtocolMsgNumbers.SSH_AGENTC_REQUEST_IDENTITIES.to_blob() r = memoryview(bytearray(self.send(req))) if r[0] != ProtocolMsgNumbers.SSH_AGENT_IDENTITIES_ANSWER: raise SshAgentFailure('agent: non-identities answer received for identities list') - return KeyList.from_blob(r[1:]) + + blob = r[1:] + nkeys, blob = uint32.consume_from_blob(blob) + rv = [] + for i in range(nkeys): + key_blob, blob = self._consume_field(blob) + + peek_key_algo, dummy = self._consume_field(key_blob) + pub_key_msg_cls = PublicKeyMsg.get_dataclass(KeyAlgo(bytes(peek_key_algo).decode('utf-8'))) + + comment_blob, blob = self._consume_field(blob) + + rv.append(Identity(pub_key_msg_cls.from_blob(key_blob), unicode_string.from_blob(comment_blob))) + + if blob: + raise SshAgentFailure("agent: received more keys than advertised") + + return rv + + @staticmethod + def _consume_field(blob: memoryview | bytes) -> tuple[memoryview | bytes, memoryview | bytes]: + length = uint32.from_blob(blob[:4]) + blob = blob[4:] + data, rest = _split_blob(blob, length) + return data, rest def __contains__(self, public_key: CryptoPublicKey) -> bool: msg = PublicKeyMsg.from_public_key(public_key) - return msg in self.list().keys + return any(i.key == msg for i in self.list()) @functools.cache diff --git a/test/integration/targets/ssh_agent/action_plugins/ssh_agent.py b/test/integration/targets/ssh_agent/action_plugins/ssh_agent.py index 979309eab25..9451a1f93a5 100644 --- a/test/integration/targets/ssh_agent/action_plugins/ssh_agent.py +++ b/test/integration/targets/ssh_agent/action_plugins/ssh_agent.py @@ -31,9 +31,9 @@ class ActionModule(ActionBase): def remove_all(self): with SshAgentClient(os.environ['SSH_AUTH_SOCK']) as client: - nkeys_before = client.list().nkeys + nkeys_before = len(client.list()) client.remove_all() - nkeys_after = client.list().nkeys + nkeys_after = len(client.list()) return { 'failed': nkeys_after != 0, 'nkeys_removed': nkeys_before, @@ -42,18 +42,18 @@ class ActionModule(ActionBase): def list(self): result = {'keys': [], 'nkeys': 0} with SshAgentClient(os.environ['SSH_AUTH_SOCK']) as client: - key_list = client.list() - result['nkeys'] = key_list.nkeys - for key in key_list.keys: - public_key = key.public_key + identity_list = client.list() + result['nkeys'] = len(identity_list) + for identity in identity_list: + public_key = identity.key.public_key key_size = getattr(public_key, 'key_size', 256) - fingerprint = key.fingerprint - key_type = key.type.main_type + fingerprint = identity.key.fingerprint + key_type = identity.key.type.main_type result['keys'].append({ 'type': key_type, 'key_size': key_size, 'fingerprint': f'SHA256:{fingerprint}', - 'comments': key.comments, + 'comment': identity.comment, }) return result