You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何优化Swift中用于RSA实现的GiantUInt大整数运算性能

问题描述

我正尝试为CryptoSwift库使用Swift实现RSA算法(用于修复#63问题),目前算法本身可以正常运行,但需要提升大整数运算的性能才能让它在合理时间内完成运算。

我自行实现了GiantUInt结构体(以UInt8格式存储字节)用于处理RSA场景下的大整数(例如2048位长度)运算,但目前运算速度过慢,主要瓶颈是取余运算,其余运算也有优化空间,代码如下:

precedencegroup PowerPrecedence { higherThan: MultiplicationPrecedence }
infix operator ^^ : PowerPrecedence

public struct GiantUInt: Equatable, Comparable, ExpressibleByIntegerLiteral, ExpressibleByArrayLiteral {
  
  // Properties
  
  public let bytes: Array<UInt8>
  
  // Initialization
  
  public init(_ raw: Array<UInt8>) {
    var bytes = raw
    
    while bytes.last == 0 {
      bytes.removeLast()
    }
    
    self.bytes = bytes
  }
  
  // ExpressibleByIntegerLiteral
  
  public typealias IntegerLiteralType = UInt8
  
  public init(integerLiteral value: UInt8) {
    self = GiantUInt([value])
  }
  
  // ExpressibleByArrayLiteral
  
  public typealias ArrayLiteralElement = UInt8
  
  public init(arrayLiteral elements: UInt8...) {
    self = GiantUInt(elements)
  }
    
  // Equatable
  
  public static func == (lhs: GiantUInt, rhs: GiantUInt) -> Bool {
    lhs.bytes == rhs.bytes
  }
  
  // Comparable
  
  public static func < (rhs: GiantUInt, lhs: GiantUInt) -> Bool {
    for i in (0 ..< max(rhs.bytes.count, lhs.bytes.count)).reversed() {
      let r = rhs.bytes[safe: i] ?? 0
      let l = lhs.bytes[safe: i] ?? 0
      if r < l {
        return true
      } else if r > l {
        return false
      }
    }
    
    return false
  }
  
  // Operations
  
  public static func + (rhs: GiantUInt, lhs: GiantUInt) -> GiantUInt {
    var bytes = [UInt8]()
    var r: UInt8 = 0
    
    for i in 0 ..< max(rhs.bytes.count, lhs.bytes.count) {
      let res = UInt16(rhs.bytes[safe: i] ?? 0) + UInt16(lhs.bytes[safe: i] ?? 0) + UInt16(r)
      r = UInt8(res >> 8)
      bytes.append(UInt8(res & 0xff))
    }
    
    if r != 0 {
      bytes.append(r)
    }
    
    return GiantUInt(bytes)
  }
  
  public static func - (rhs: GiantUInt, lhs: GiantUInt) -> GiantUInt {
    var bytes = [UInt8]()
    var r: UInt8 = 0
    
    for i in 0 ..< max(rhs.bytes.count, lhs.bytes.count) {
      let rhsb = UInt16(rhs.bytes[safe: i] ?? 0)
      let lhsb = UInt16(lhs.bytes[safe: i] ?? 0) + UInt16(r)
      r = UInt8(rhsb < lhsb ? 1 : 0)
      let res = (UInt16(r) << 8) + rhsb - lhsb
      bytes.append(UInt8(res & 0xff))
    }
    
    if r != 0 {
      bytes.append(r)
    }
    
    return GiantUInt(bytes)
  }
  
  public static func * (rhs: GiantUInt, lhs: GiantUInt) -> GiantUInt {
    var offset = 0
    var sum = [GiantUInt]()
    
    for rbyte in rhs.bytes {
      var bytes = [UInt8](repeating: 0, count: offset)
      var r: UInt8 = 0
      
      for lbyte in lhs.bytes {
        let res = UInt16(rbyte) * UInt16(lbyte) + UInt16(r)
        r = UInt8(res >> 8)
        bytes.append(UInt8(res & 0xff))
      }
      
      if r != 0 {
        bytes.append(r)
      }
      
      sum.append(GiantUInt(bytes))
      offset += 1
    }
    
    return sum.reduce(0, +)
  }
  
  public static func % (rhs: GiantUInt, lhs: GiantUInt) -> GiantUInt {
    var remainder = rhs
    
    // This needs serious optimization (but works)
    while remainder >= lhs {
      remainder = remainder - lhs
    }
  
    return remainder
  }
  
  static func ^^ (rhs: GiantUInt, lhs: GiantUInt) -> GiantUInt {
    let count = lhs.bytes.count
    var result = GiantUInt([1])
    
    for iByte in 0 ..< count {
      let byte = lhs.bytes[iByte]
      for i in 0 ..< 8 {
        if iByte != count - 1 || byte >> i > 0 {
          result = result * result
          if (byte >> i) & 1 == 1 {
            result = result * rhs
          }
        }
      }
    }
    
    return result
  }
  
  public static func exponentiateWithModulus(rhs: GiantUInt, lhs: GiantUInt, modulus: GiantUInt) -> GiantUInt {
    let count = lhs.bytes.count
    var result = GiantUInt([1])
    
    for iByte in 0 ..< count {
      let byte = lhs.bytes[iByte]
      for i in 0 ..< 8 {
        if iByte != count - 1 || byte >> i > 0 {
          result = (result * result) % modulus
          if (byte >> i) & 1 == 1 {
            result = (result * rhs) % modulus
          }
        }
      }
    }
    
    return result
  }
  
}

请问如何优化上述实现才能大幅提升运算速度?


优化方案
  • 核心瓶颈优先修复:替换取余实现
    你当前的取余逻辑是反复减去除数,对于2048位的大整数来说,运算次数会达到数十亿次,完全不可用。请直接替换为二进制长除法实现的取余,单次取余时间复杂度可以降到O(nm),性能至少提升几个数量级。如果是专门用于RSA的模幂场景,直接实现蒙哥马利约减*算法,可以进一步大幅降低连续模运算的开销。
  • 调整存储单元提升基础运算效率
    当前用UInt8作为最小存储单元太低效,现代CPU都是64位架构,建议把存储单元替换为UInt64(通常称为limb),单次运算可以处理64位数据,相当于把原来需要8次的UInt8运算合并为1次,直接带来数倍的基础运算性能提升,还可以充分利用CPU原生的64位运算指令。
  • 优化乘法运算实现
    当前的乘法实现会生成大量临时GiantUInt对象再累加,内存开销和运算开销都很高。建议直接维护一个结果数组,每次乘法的中间结果直接累加到数组中,避免临时对象的创建和拷贝。对于长度超过256位的大整数,可以实现Karatsuba乘法,把乘法的时间复杂度从O(n²)降到O(n^1.585),位数越高提升越明显。
  • 模幂运算针对性优化
    你当前的快速幂逻辑已经可用,结合蒙哥马利约减之后,可以把所有中间运算都放在蒙哥马利域下完成,最后再转换为普通整数,模运算的开销可以降低一半以上。还可以采用滑动窗口法优化快速幂的执行流程,减少不必要的乘法操作。
  • 通用细节优化
    移除所有不必要的安全下标访问,在运算逻辑可控的前提下直接用原生下标访问,省去边界判断的开销。把GiantUInt改为可变结构体,运算时直接修改内部存储数组,避免每次运算都创建新的数组和结构体对象,减少内存分配和拷贝的开销。运算前先判断操作数长度,针对短操作数走快速路径,避免不必要的循环。

内容的提问来源于stack exchange,提问作者Nathan Fallet

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.06 20:09:01