mirror of
https://github.com/ansible/ansible.git
synced 2026-07-28 08:05:22 +02:00
ssh-agent: hardening (#86921)
* fix partial socket reads * cap agent response size to 256 KiB * fix wire format for zero-length binary_string, unicode_string and mpint * separate key and comment in identities response per the protocol * remove KeyList/PublicKeyMsgList in favor of inline parsing and Identity dataclass * rename comments -> comment * switch KeyAlgo to StrEnum ci_complete
This commit is contained in:
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user