Skip to content

Commit a0795a6

Browse files
committed
use der crate for parsing
1 parent a2b19da commit a0795a6

1 file changed

Lines changed: 97 additions & 192 deletions

File tree

russh/src/keys/format/pkcs8.rs

Lines changed: 97 additions & 192 deletions
Original file line numberDiff line numberDiff line change
@@ -24,28 +24,22 @@ pub fn decode_pkcs8(
2424
doc
2525
};
2626

27-
match doc.decode_msg::<sec1::EcPrivateKey>() {
28-
Ok(key) => {
29-
// X9.62 EC private key
30-
let Some(curve) = key.parameters.and_then(|x| x.named_curve()) else {
31-
return Err(Error::CouldNotReadKey);
32-
};
33-
let kp = ec_key_data_into_keypair(curve, key)?;
34-
Ok(PrivateKey::new(KeypairData::Ecdsa(kp), "")?)
35-
}
36-
Err(_) => {
37-
// SEC1 key with full domain parameters (not a named curve OID)
38-
match decode_sec1_with_full_domain_params(ciphertext) {
39-
Ok(kp) => return Ok(PrivateKey::new(KeypairData::Ecdsa(kp), "")?),
40-
Err(_) => {},
41-
}
42-
// ASN.1 key (PKCS#8)
43-
Ok(
44-
pkcs8_pki_into_keypair_data(doc.decode_msg::<PrivateKeyInfoRef<'_>>()?)?
45-
.try_into()?,
46-
)
47-
}
27+
if let Ok(key) = doc.decode_msg::<sec1::EcPrivateKey>() {
28+
// X9.62 EC private key
29+
let Some(curve) = key.parameters.and_then(|x| x.named_curve()) else {
30+
return Err(Error::CouldNotReadKey);
31+
};
32+
let kp = ec_key_data_into_keypair(curve, key)?;
33+
return Ok(PrivateKey::new(KeypairData::Ecdsa(kp), "")?);
34+
}
35+
36+
// SEC1 key with full domain parameters (not a named curve OID)
37+
if let Ok(kp) = explicit_curve_params::decode_sec1_with_full_domain_params(ciphertext) {
38+
return Ok(PrivateKey::new(KeypairData::Ecdsa(kp), "")?);
4839
}
40+
41+
// ASN.1 key (PKCS#8)
42+
Ok(pkcs8_pki_into_keypair_data(doc.decode_msg::<PrivateKeyInfoRef<'_>>()?)?.try_into()?)
4943
}
5044

5145
fn pkcs8_pki_into_keypair_data(pki: PrivateKeyInfoRef<'_>) -> Result<KeypairData, Error> {
@@ -117,191 +111,102 @@ where
117111
})
118112
}
119113

120-
/// Try to manually parse an SEC1 EC key with full domain parameters.
121-
///
122-
/// Some key generators (e.g. OpenSSL with certain options) produce SEC1 keys
123-
/// where the `[0]` parameters field contains full EC domain parameters instead
124-
/// of a named curve OID. The `sec1` crate does not support this format.
125-
/// This function manually parses the DER to extract the curve and private key.
126-
fn decode_sec1_with_full_domain_params(der: &[u8]) -> Result<EcdsaKeypair, Error> {
127-
if der.is_empty() || der[0] != 0x30 {
128-
return Err(Error::CouldNotReadKey);
129-
}
130-
let (header_size, _) = read_der_len(&der[1..], der.len() - 1)?;
131-
let mut pos = 1 + header_size;
114+
mod explicit_curve_params {
115+
use super::*;
132116

133-
if pos >= der.len() || der[pos] != 0x02 {
134-
return Err(Error::CouldNotReadKey);
135-
}
136-
let (ver, new_pos) = read_der_value(der, pos)?;
137-
pos = new_pos;
138-
// SEC1 version must be 1 (ecPrivkeyVer1)
139-
if ver.len() != 1 || ver[0] < 1 {
140-
return Err(Error::CouldNotReadKey);
141-
}
117+
use der::{
118+
Reader, SliceReader, Tag, TagNumber, Tagged,
119+
asn1::{AnyRef, ContextSpecific, UintRef},
120+
};
142121

143-
if pos >= der.len() || der[pos] != 0x04 {
144-
return Err(Error::CouldNotReadKey);
145-
}
146-
let (priv_key, new_pos) = read_der_value(der, pos)?;
147-
pos = new_pos;
122+
/// Try to parse an SEC1 EC key with full domain parameters.
123+
///
124+
/// Some key generators (e.g. OpenSSL with certain options) produce SEC1 keys
125+
/// where the `[0]` parameters field contains full EC domain parameters instead
126+
/// of a named curve OID. The `sec1` crate does not support this format.
127+
pub fn decode_sec1_with_full_domain_params(der_bytes: &[u8]) -> Result<EcdsaKeypair, Error> {
128+
let mut reader = SliceReader::new(der_bytes)?;
129+
reader.sequence(|seq| {
130+
let version: u8 = seq.decode()?;
131+
if version < 1 {
132+
return Err(Error::CouldNotReadKey);
133+
}
148134

149-
if pos >= der.len() || der[pos] != 0xa0 {
150-
return Err(Error::CouldNotReadKey);
151-
}
152-
let (params_content, _) = read_der_value(der, pos)?;
135+
let priv_key: AnyRef = seq.decode()?;
136+
priv_key.tag().assert_eq(Tag::OctetString)?;
153137

154-
if params_content.is_empty() || params_content[0] != 0x30 {
155-
return Err(Error::CouldNotReadKey);
156-
}
138+
let params = ContextSpecific::<AnyRef>::decode_explicit(seq, TagNumber(0))?
139+
.ok_or(Error::CouldNotReadKey)?;
157140

158-
let curve_oid = extract_curve_from_domain_params(params_content)?;
159-
build_ec_keypair_from_bytes(curve_oid, priv_key)
160-
}
141+
let curve_oid = extract_curve_from_domain_params(params.value)?;
161142

162-
/// Read a DER TLV (tag-length-value) from `der` starting at `pos`.
163-
/// Returns (value_bytes, next_position_after_this_tlv).
164-
fn read_der_value(der: &[u8], pos: usize) -> Result<(&[u8], usize), Error> {
165-
if pos >= der.len() {
166-
return Err(Error::CouldNotReadKey);
167-
}
168-
// read_der_len reads from the LENGTH field (pos + 1, after the tag byte).
169-
// It returns (length_field_size, length_field_size + content_length).
170-
// So total_size = header_size + content_size (relative to length field start).
171-
let (length_field_size, total_from_len_field) = read_der_len(&der[pos + 1..], der.len() - pos - 1)?;
172-
let header_size = 1 + length_field_size; // tag byte + length field bytes
173-
let value_start = pos + header_size;
174-
let value_end = pos + 1 + total_from_len_field; // 1 for tag + total_from_len_field
175-
if value_end > der.len() || value_start > value_end {
176-
return Err(Error::CouldNotReadKey);
177-
}
178-
Ok((&der[value_start..value_end], value_end))
179-
}
143+
let keypair = build_ec_keypair_from_bytes(curve_oid, priv_key.value())?;
180144

181-
/// Read a DER length field from `buf` (which starts at the length byte, after the tag).
182-
///
183-
/// Returns `(length_field_size, length_field_size + content_length)`:
184-
/// - `length_field_size`: number of bytes used to encode the length itself (1 for short form, 1+N for long form)
185-
/// - `length_field_size + content_length`: total bytes consumed (length field + content)
186-
///
187-
/// `max` is the number of available bytes in `buf` (used for bounds checking).
188-
fn read_der_len(buf: &[u8], max: usize) -> Result<(usize, usize), Error> {
189-
if buf.is_empty() {
190-
return Err(Error::CouldNotReadKey);
191-
}
192-
let first = buf[0];
193-
if first & 0x80 == 0 {
194-
let len = first as usize;
195-
if 1 + len > max {
196-
return Err(Error::CouldNotReadKey);
197-
}
198-
Ok((1, 1 + len))
199-
} else {
200-
let num_bytes = (first & 0x7f) as usize;
201-
if num_bytes == 0 || 1 + num_bytes > max {
202-
return Err(Error::CouldNotReadKey);
203-
}
204-
let mut len = 0usize;
205-
for &b in &buf[1..1 + num_bytes] {
206-
len = (len << 8) | b as usize;
207-
}
208-
if 1 + num_bytes + len > max {
209-
return Err(Error::CouldNotReadKey);
210-
}
211-
Ok((1 + num_bytes, 1 + num_bytes + len))
212-
}
213-
}
214-
215-
/// Extract the named curve OID from full EC domain parameters.
216-
/// Handles two formats:
217-
/// 1. Standard ECParameters: SEQUENCE { FieldID, Curve, base, order, cofactor }
218-
/// 2. Wrapped ECParameters: SEQUENCE { INTEGER version, SEQUENCE { FieldID, ... } }
219-
fn extract_curve_from_domain_params(params_der: &[u8]) -> Result<ObjectIdentifier, Error> {
220-
if params_der.is_empty() || params_der[0] != 0x30 {
221-
return Err(Error::CouldNotReadKey);
222-
}
223-
let (header_size, _) = read_der_len(&params_der[1..], params_der.len() - 1)?;
224-
let mut pos = 1 + header_size; // skip SEQUENCE tag + length bytes
225-
226-
if pos >= params_der.len() {
227-
return Err(Error::CouldNotReadKey);
145+
// Drain any remaining optional fields (e.g. [1] publicKey) so finish() succeeds
146+
seq.drain(seq.remaining_len())?;
147+
Ok(keypair)
148+
})
228149
}
229150

230-
// If the first element is an INTEGER (version), skip it to reach FieldID SEQUENCE
231-
if params_der[pos] == 0x02 {
232-
let (_, skip_pos) = read_der_value(params_der, pos)?;
233-
pos = skip_pos;
234-
}
151+
/// Extract the named curve OID from full EC domain parameters.
152+
/// Handles two formats:
153+
/// 1. Standard ECParameters: SEQUENCE { FieldID, Curve, base, order, cofactor }
154+
/// 2. Wrapped ECParameters: SEQUENCE { INTEGER version, SEQUENCE { FieldID, ... } }
155+
fn extract_curve_from_domain_params(params: AnyRef<'_>) -> Result<ObjectIdentifier, Error> {
156+
params.tag().assert_eq(Tag::Sequence)?;
235157

236-
// Now pos should point to the FieldID SEQUENCE
237-
if pos >= params_der.len() || params_der[pos] != 0x30 {
238-
return Err(Error::CouldNotReadKey);
239-
}
240-
let (field_id_content, _new_pos) = read_der_value(params_der, pos)?;
158+
// Use a standalone SliceReader so we aren't required to consume all of ECParams
159+
// (Curve, base, order, cofactor follow FieldID but are irrelevant here).
160+
let mut seq = SliceReader::new(params.value())?;
241161

242-
// Parse FieldID: first element is OID (prime-field = 1.2.840.10045.1.1)
243-
if field_id_content.is_empty() || field_id_content[0] != 0x06 {
244-
return Err(Error::CouldNotReadKey);
245-
}
246-
// OID 1.2.840.10045.1.1 (prime-field from ANSI X9.62):
247-
// DER tag 0x06, length 0x07, then 7 bytes of OID content
248-
// 2a = 1.2, 86 48 = 840, ce 3d = 10045, 01 = 1, 01 = 1
249-
// This is a stable ASN.1 standard OID that will not change.
250-
let prime_field_der: &[u8] = &[0x06, 0x07, 0x2a, 0x86, 0x48, 0xce, 0x3d, 0x01, 0x01];
251-
let oid_len = field_id_content[1] as usize;
252-
let oid_end = 2 + oid_len;
253-
if oid_end > field_id_content.len() || &field_id_content[..oid_end] != prime_field_der {
254-
return Err(Error::CouldNotReadKey);
255-
}
256-
let prime_pos = oid_end;
162+
// Skip optional ECParameters version INTEGER
163+
if Tag::peek(&seq)? == Tag::Integer {
164+
seq.decode::<u8>()?;
165+
}
257166

258-
if prime_pos >= field_id_content.len() || field_id_content[prime_pos] != 0x02 {
259-
return Err(Error::CouldNotReadKey);
167+
// FieldID ::= SEQUENCE { fieldType OID, parameters ANY }
168+
seq.sequence(|field_id| {
169+
let _field_oid: ObjectIdentifier = field_id.decode()?;
170+
// prime INTEGER — as_bytes() strips DER sign-extension leading zero
171+
let prime: UintRef = field_id.decode()?;
172+
Ok(match prime.as_bytes().len() {
173+
32 => NistP256::OID,
174+
48 => NistP384::OID,
175+
66 => NistP521::OID,
176+
_ => return Err(Error::CouldNotReadKey),
177+
})
178+
})
260179
}
261-
let (prime_bytes, _) = read_der_value(field_id_content, prime_pos)?;
262-
263-
// Determine curve from prime byte length.
264-
// DER INTEGERs are signed, so when the high bit of the first content byte is set,
265-
// a leading 0x00 pad byte is added to keep the value positive. For example, the
266-
// P-256 prime (a 256-bit value with high bit set) encodes as 33 bytes: 0x00 || 32-byte-prime.
267-
// We match both padded (33/49/67) and unpadded (32/48/66) lengths.
268-
let prime_len = prime_bytes.len();
269-
Ok(match prime_len {
270-
32 | 33 => NistP256::OID, // P-256: secp256r1 prime is 32 bytes
271-
48 | 49 => NistP384::OID, // P-384: secp384r1 prime is 48 bytes
272-
66 | 67 => NistP521::OID, // P-521: secp521r1 prime is 66 bytes
273-
_ => return Err(Error::CouldNotReadKey),
274-
})
275-
}
276180

277-
/// Build an EcdsaKeypair from raw private key bytes and a curve OID.
278-
fn build_ec_keypair_from_bytes(
279-
curve_oid: ObjectIdentifier,
280-
private_key_bytes: &[u8],
281-
) -> Result<EcdsaKeypair, Error> {
282-
if curve_oid == NistP256::OID {
283-
let sk = p256::SecretKey::from_slice(private_key_bytes)
284-
.map_err(|_| Error::CouldNotReadKey)?;
285-
Ok(EcdsaKeypair::NistP256 {
286-
public: sk.public_key().into(),
287-
private: sk.into(),
288-
})
289-
} else if curve_oid == NistP384::OID {
290-
let sk = p384::SecretKey::from_slice(private_key_bytes)
291-
.map_err(|_| Error::CouldNotReadKey)?;
292-
Ok(EcdsaKeypair::NistP384 {
293-
public: sk.public_key().into(),
294-
private: sk.into(),
295-
})
296-
} else if curve_oid == NistP521::OID {
297-
let sk = p521::SecretKey::from_slice(private_key_bytes)
298-
.map_err(|_| Error::CouldNotReadKey)?;
299-
Ok(EcdsaKeypair::NistP521 {
300-
public: sk.public_key().into(),
301-
private: sk.into(),
302-
})
303-
} else {
304-
Err(Error::UnknownAlgorithm(curve_oid))
181+
/// Build an EcdsaKeypair from raw private key bytes and a curve OID.
182+
fn build_ec_keypair_from_bytes(
183+
curve_oid: ObjectIdentifier,
184+
private_key_bytes: &[u8],
185+
) -> Result<EcdsaKeypair, Error> {
186+
if curve_oid == NistP256::OID {
187+
let sk = p256::SecretKey::from_slice(private_key_bytes)
188+
.map_err(|_| Error::CouldNotReadKey)?;
189+
Ok(EcdsaKeypair::NistP256 {
190+
public: sk.public_key().into(),
191+
private: sk.into(),
192+
})
193+
} else if curve_oid == NistP384::OID {
194+
let sk = p384::SecretKey::from_slice(private_key_bytes)
195+
.map_err(|_| Error::CouldNotReadKey)?;
196+
Ok(EcdsaKeypair::NistP384 {
197+
public: sk.public_key().into(),
198+
private: sk.into(),
199+
})
200+
} else if curve_oid == NistP521::OID {
201+
let sk = p521::SecretKey::from_slice(private_key_bytes)
202+
.map_err(|_| Error::CouldNotReadKey)?;
203+
Ok(EcdsaKeypair::NistP521 {
204+
public: sk.public_key().into(),
205+
private: sk.into(),
206+
})
207+
} else {
208+
Err(Error::UnknownAlgorithm(curve_oid))
209+
}
305210
}
306211
}
307212

0 commit comments

Comments
 (0)