dApp Docs/AI Agent 联邦学习与隐私计算指南
Development reference. Not independently verified for production.

MSG Chain AI Agent 联邦学习与隐私计算指南

适用链: msg-chain-1 | 地址前缀: msg | 版本: v1.0


目录

  1. 概述
  2. 联邦学习架构
  3. 梯度聚合合约
  4. 差分隐私保护
  5. TEE 机密计算集成
  6. ZKP 推理证明
  7. 模型激励与贡献度量
  8. 联邦学习市场 App

1. 概述

1.1 为什么 AI Agent 需要联邦学习

在 MSG Chain 上运行的 AI Agent 面临一个根本性矛盾:模型需要更多数据来提升智能,但数据主权和隐私法规(如 GDPR、个保法)禁止原始数据共享。联邦学习(Federated Learning, FL)为这一矛盾提供了解决方案。

传统中心化机器学习的工作流如下:

所有 Agent 上传原始数据 → 中心服务器训练 → 分发最终模型

联邦学习翻转了这一流程:

中心服务器分发初始模型 → Agent 本地训练 → 仅上传梯度更新 → 服务器聚合

对于 MSG Chain 上的 AI Agent 生态,联邦学习带来以下核心价值:

1.2 Agent 数据的隐私挑战

MSG Chain 上的 AI Agent 处理的数据类型包括:

数据类型 示例 隐私风险
对话历史 用户与 Agent 的聊天记录 包含 PII、商业机密
交易行为 DeFi 交互、NFT 交易 可关联身份与金融活动
个人偏好 Agent 个性化配置 用户画像泄露
知识库 RAG 索引的私有文档 知识产权泄露

即使在联邦学习中,仅上传梯度更新而非原始数据,仍然存在隐私泄露风险:

  1. 梯度泄露攻击(Gradient Leakage):攻击者可以从梯度中重建原始训练数据。研究表明,在图像和文本任务中,通过优化噪声梯度可以高保真度还原输入样本。
  2. 成员推断攻击(Membership Inference):通过观察模型更新,攻击者可以推断某个特定样本是否在 Agent 的训练数据中。
  3. 模型逆向(Model Inversion):从模型参数中重建训练数据的统计特征。

1.3 TEE + ZKP + FL 三位一体防护

为应对上述挑战,本指南采用 三层防护架构:

┌─────────────────────────────────────────────┐
│  联邦学习 (FL)                               │
│  ─ 数据不出本地,仅交换梯度                    │
│  ─ SecAgg 安全聚合                            │
├─────────────────────────────────────────────┤
│  差分隐私 (DP)                                │
│  ─ 梯度加噪,抵御泄露攻击                      │
│  ─ 隐私预算追踪                               │
├─────────────────────────────────────────────┤
│  可信执行环境 (TEE) + 零知识证明 (ZKP)         │
│  ─ TEE: 硬件级隔离执行                        │
│  ─ ZKP: 可验证的推理正确性                    │
│  ─ 远程证明: 确保代码未被篡改                  │
└─────────────────────────────────────────────┘

三层协作流程:

  1. Agent 在本地 TEE 环境中训练模型,确保训练过程对操作系统和云提供商不可见。
  2. 梯度在离开 TEE 前经过 差分隐私 加噪处理,即使梯度被截获也无法还原原始数据。
  3. 通过 安全聚合(SecAgg),服务端只能看到聚合后的梯度,无法区分单个 Agent 的贡献。
  4. 推理阶段,Agent 使用 零知识证明 证明推理结果的正确性,而不泄露模型参数或输入数据。
  5. 聚合结果写入 MSG Chain 智能合约,通过 梯度聚合合约 验证并分配奖励。

2. 联邦学习架构

2.1 架构总览

联邦学习在 MSG Chain Agent 生态中有两种部署模式:中心化 FL 和 去中心化 FL。

2.1.1 中心化 FL(Client-Server)

        ┌──────────────────┐
        │  FL Coordinator  │  ← MSG Chain 合约
        │  (智能合约)       │
        └────┬──────┬──────┘
             │      │
     ┌───────▼┐  ┌─▼───────┐
     │ Agent A│  │ Agent B │  ← TEE 环境
     │ 本地训练│  │ 本地训练 │
     └────────┘  └─────────┘

适用场景:Agent 数量少(<100),需要严格协调训练轮次。

2.1.2 去中心化 FL(Gossip 协议)

    ┌──────────┐     ┌──────────┐
    │ Agent A  │◄───►│ Agent B  │
    │ 模型 v1  │     │ 模型 v1  │
    └────┬─────┘     └────┬─────┘
         │                │
    ┌────▼─────┐     ┌────▼─────┐
    │ Agent C  │◄───►│ Agent D  │
    │ 模型 v1  │     │ 模型 v1  │
    └──────────┘     └──────────┘

适用场景:大规模 Agent 网络(>100),需要高容错性和可扩展性。

2.2 联邦学习协调器

以下实现了一个完整的联邦学习协调器,运行在 MSG Chain Agent 环境中:

import asyncio
import logging
import time
from dataclasses import dataclass, field
from typing import Dict, List, Optional, Tuple, Any
from enum import Enum

import numpy as np
import msg_client  # MSG Chain Python SDK

logger = logging.getLogger(__name__)


class AgentStatus(Enum):
    IDLE = "idle"
    TRAINING = "training"
    UPLOADING = "uploading"
    VERIFIED = "verified"
    SLASHED = "slashed"


@dataclass
class AgentInfo:
    """Agent 在联邦学习中的注册信息"""
    address: str
    stake: int
    reputation: float
    data_quality: float
    status: AgentStatus = AgentStatus.IDLE
    last_round: int = 0
    training_time_ms: int = 0
    contribution_score: float = 0.0


@dataclass
class ModelMetrics:
    """模型训练指标"""
    accuracy: float = 0.0
    loss: float = float('inf')
    num_samples: int = 0
    round_time_ms: int = 0


class Model:
    """简化的神经网络模型表示"""

    def __init__(self, layers: Optional[List[Dict]] = None):
        self.layers = layers or [
            {"type": "dense", "units": 128, "activation": "relu"},
            {"type": "dense", "units": 64, "activation": "relu"},
            {"type": "dense", "units": 10, "activation": "softmax"},
        ]
        self.weights: List[np.ndarray] = []
        self.biases: List[np.ndarray] = []
        self.version: int = 0
        self.merkle_root: str = ""

    def serialize(self) -> bytes:
        """序列化模型权重为字节流"""
        data = b""
        for w, b in zip(self.weights, self.biases):
            data += w.tobytes() + b.tobytes()
        return data

    def deserialize(self, data: bytes, shapes: List[Tuple[int, ...]]):
        """从字节流反序列化模型权重"""
        offset = 0
        self.weights = []
        self.biases = []
        for shape_w, shape_b in shapes:
            w_size = int(np.prod(shape_w))
            b_size = int(np.prod(shape_b))
            w = np.frombuffer(data[offset:offset + w_size * 4], dtype=np.float32).reshape(shape_w)
            offset += w_size * 4
            b = np.frombuffer(data[offset:offset + b_size * 4], dtype=np.float32).reshape(shape_b)
            offset += b_size * 4
            self.weights.append(w)
            self.biases.append(b)

    def commitment(self) -> str:
        """计算模型承诺(Merkle Root)"""
        import hashlib
        h = hashlib.sha256()
        for w in self.weights:
            h.update(w.tobytes())
        for b in self.biases:
            h.update(b.tobytes())
        return h.hexdigest()

    def copy(self) -> "Model":
        new_model = Model(self.layers.copy())
        new_model.weights = [w.copy() for w in self.weights]
        new_model.biases = [b.copy() for b in self.biases]
        new_model.version = self.version
        return new_model


@dataclass
class TrainingRound:
    """单轮训练的状态"""
    round_id: int
    coordinator: str
    global_model: Model
    participants: List[str] = field(default_factory=list)
    selected_agents: List[str] = field(default_factory=list)
    gradients: Dict[str, bytes] = field(default_factory=dict)
    metrics: Dict[str, ModelMetrics] = field(default_factory=dict)
    start_time: float = 0.0
    deadline: float = 0.0
    min_participants: int = 3
    aggregation_weights: Dict[str, float] = field(default_factory=dict)


class FederatedLearningCoordinator:
    """跨 Agent 协调联邦学习的核心类

    负责 Agent 注册、轮次管理、安全聚合、奖励分配。
    与 MSG Chain 智能合约交互以实现链上协调。
    """

    def __init__(
        self,
        registry_client: msg_client.Client,
        min_agents: int = 3,
        round_timeout: int = 300,
        aggregation_algorithm: str = "fedavg",
    ):
        self.client = registry_client
        self.agents: Dict[str, AgentInfo] = {}
        self.global_model = Model()
        self.round = 0
        self.min_agents = min_agents
        self.round_timeout = round_timeout
        self.aggregation_algorithm = aggregation_algorithm
        self.training_history: List[TrainingRound] = []
        self.privacy_budget_per_round: float = 0.1
        self.active_round: Optional[TrainingRound] = None

    async def register_agent(
        self,
        agent_address: str,
        stake: int,
        data_quality_score: float = 0.5,
    ) -> bool:
        """在 FL 网络中注册 Agent

        Agent 需要质押代币以参与训练,质押金额影响其贡献权重。
        """
        if agent_address in self.agents:
            logger.warning(f"Agent {agent_address} 已注册")
            return False

        if stake < 100:
            logger.warning(f"Agent {agent_address} 质押不足,需要至少 100 MSG")
            return False

        agent = AgentInfo(
            address=agent_address,
            stake=stake,
            reputation=0.5,
            data_quality=data_quality_score,
        )
        self.agents[agent_address] = agent

        # 链上注册事件
        await self.client.send_tx(
            contract="msg1flregistry...",
            action="register_agent",
            params={
                "agent": agent_address,
                "stake": stake,
                "quality_score": data_quality_score,
            },
        )
        logger.info(f"Agent {agent_address} 注册成功,质押 {stake} MSG")
        return True

    async def unregister_agent(self, agent_address: str) -> bool:
        """注销 Agent"""
        if agent_address not in self.agents:
            return False
        del self.agents[agent_address]

        await self.client.send_tx(
            contract="msg1flregistry...",
            action="unregister_agent",
            params={"agent": agent_address},
        )
        return True

    async def select_agents_for_round(
        self,
        round_id: int,
        max_agents: int = 10,
        min_reputation: float = 0.3,
    ) -> List[str]:
        """选择参与本轮训练的 Agent

        选择策略基于:
        1. 质押金额(越高越优先)
        2. 信誉分数(越高越优先)
        3. 数据质量评分
        4. 上次参与时间(优先选择长时间未参与的 Agent)
        5. 随机性(防止固定选择偏差)
        """
        eligible = [
            addr for addr, info in self.agents.items()
            if info.status == AgentStatus.IDLE
            and info.reputation >= min_reputation
        ]

        if len(eligible) < self.min_agents:
            raise RuntimeError(
                f"合格 Agent 不足: {len(eligible)} < {self.min_agents}"
            )

        # 综合评分选择
        scores = {}
        for addr in eligible:
            info = self.agents[addr]
            base_score = (
                info.stake * 0.3
                + info.reputation * 0.4
                + info.data_quality * 0.2
            )
            # 时间衰减因子:长时间未参与的 Agent 获得加分
            rounds_since_participation = round_id - info.last_round
            time_bonus = min(rounds_since_participation * 0.05, 0.5)
            # 随机扰动
            random_factor = np.random.uniform(0.9, 1.1)
            final_score = (base_score + time_bonus) * random_factor
            scores[addr] = final_score

        # 按分数排序
        sorted_agents = sorted(scores.items(), key=lambda x: x[1], reverse=True)
        selected = [addr for addr, _ in sorted_agents[:max_agents]]

        # 更新状态
        for addr in selected:
            self.agents[addr].status = AgentStatus.TRAINING
            self.agents[addr].last_round = round_id

        logger.info(
            f"第 {round_id} 轮选择了 {len(selected)} 个 Agent: {selected[:5]}..."
        )
        return selected

    async def distribute_global_model(
        self,
        selected_agents: List[str],
        model: Model,
    ) -> Dict[str, bool]:
        """向选中的 Agent 分发当前全局模型

        通过 MSG Chain 的跨 Agent 消息传递机制分发。
        """
        tasks = []
        for agent_addr in selected_agents:
            task = self.client.send_message(
                to=agent_addr,
                message_type="model_distribution",
                payload={
                    "model_version": model.version,
                    "model_weights": model.serialize().hex(),
                    "model_commitment": model.commitment(),
                    "round_id": self.round,
                    "deadline": int(time.time()) + self.round_timeout,
                },
            )
            tasks.append(task)

        results = await asyncio.gather(*tasks, return_exceptions=True)
        delivery_status = {}
        for i, addr in enumerate(selected_agents):
            if isinstance(results[i], Exception):
                logger.error(f"向 {addr} 分发模型失败: {results[i]}")
                delivery_status[addr] = False
            else:
                delivery_status[addr] = True

        return delivery_status

    async def collect_gradients(
        self,
        selected_agents: List[str],
        timeout: int = 300,
    ) -> Dict[str, bytes]:
        """收集 Agent 本地训练后的梯度更新

        使用安全聚合协议,服务端只能看到聚合结果。
        """
        deadline = time.time() + timeout
        gradients: Dict[str, bytes] = {}

        while time.time() < deadline:
            for addr in selected_agents:
                if addr in gradients:
                    continue
                try:
                    result = await self.client.query_message(
                        from_addr=addr,
                        message_type="gradient_upload",
                        round_id=self.round,
                    )
                    if result:
                        gradients[addr] = bytes.fromhex(result["gradient_hex"])
                        self.agents[addr].status = AgentStatus.UPLOADING
                        self.agents[addr].training_time_ms = result.get(
                            "training_time_ms", 0
                        )
                        logger.info(f"收到 {addr} 的梯度 ({len(gradients[addr])} bytes)")
                except Exception as e:
                    await asyncio.sleep(0.5)
                    continue

            if len(gradients) >= len(selected_agents):
                break
            await asyncio.sleep(1)

        # 超时未提交的标记为 slash
        for addr in selected_agents:
            if addr not in gradients:
                logger.warning(f"Agent {addr} 超时未提交梯度")
                self.agents[addr].status = AgentStatus.SLASHED
                await self._slash_agent(addr)

        return gradients

    async def _slash_agent(self, agent_address: str):
        """惩罚未履行训练义务的 Agent"""
        info = self.agents[agent_address]
        slash_amount = info.stake // 10
        info.stake -= slash_amount
        info.reputation = max(0.0, info.reputation - 0.1)

        await self.client.send_tx(
            contract="msg1flregistry...",
            action="slash_agent",
            params={
                "agent": agent_address,
                "amount": slash_amount,
                "reason": "timeout",
            },
        )

    def secure_aggregate(
        self,
        gradients: Dict[str, bytes],
        model: Model,
        weights: Optional[Dict[str, float]] = None,
    ) -> Model:
        """安全聚合 — FedAvg 算法

        使用加权平均聚合所有 Agent 的梯度更新。
        权重基于 Agent 的数据样本量和信誉分数。
        """
        if not gradients:
            return model

        # 反序列化梯度
        grad_list = []
        weight_list = []
        total_weight = 0.0

        for addr, grad_bytes in gradients.items():
            agent_weight = weights.get(addr, 1.0) if weights else 1.0
            agent_weight *= self.agents[addr].reputation
            total_weight += agent_weight
            grad_list.append(grad_bytes)
            weight_list.append(agent_weight)

        # 按权重聚合
        aggregated_model = model.copy()
        new_weights = []
        new_biases = []

        for layer_idx in range(len(model.weights)):
            w_shape = model.weights[layer_idx].shape
            b_shape = model.biases[layer_idx].shape
            w_accum = np.zeros(w_shape, dtype=np.float32)
            b_accum = np.zeros(b_shape, dtype=np.float32)

            for i, grad_bytes in enumerate(grad_list):
                offset = 0
                # 跳过所有之前的层
                for j in range(layer_idx):
                    w_s = model.weights[j].shape
                    b_s = model.biases[j].shape
                    offset += int(np.prod(w_s) * 4) + int(np.prod(b_s) * 4)

                grad = np.frombuffer(
                    grad_bytes[offset:offset + int(np.prod(w_shape) * 4)],
                    dtype=np.float32,
                ).reshape(w_shape)
                bias = np.frombuffer(
                    grad_bytes[
                        offset + int(np.prod(w_shape) * 4):
                        offset + int(np.prod(w_shape) * 4) + int(np.prod(b_shape) * 4)
                    ],
                    dtype=np.float32,
                ).reshape(b_shape)

                w_accum += grad * (weight_list[i] / total_weight)
                b_accum += bias * (weight_list[i] / total_weight)

            new_weights.append(w_accum)
            new_biases.append(b_accum)

        aggregated_model.weights = new_weights
        aggregated_model.biases = new_biases
        aggregated_model.version = model.version + 1

        return aggregated_model

    async def start_training_round(self) -> TrainingRound:
        """启动一轮完整的联邦学习训练

        流程:
        1. 选择参与 Agent
        2. 分发全局模型
        3. 收集本地梯度更新
        4. 安全聚合
        5. 更新全局模型
        6. 记录训练历史
        7. 分发奖励
        """
        self.round += 1
        round_id = self.round
        logger.info(f"=== 开始第 {round_id} 轮联邦学习 ===")

        # 记录轮次
        training_round = TrainingRound(
            round_id=round_id,
            coordinator=self.client.address,
            global_model=self.global_model.copy(),
            start_time=time.time(),
            deadline=time.time() + self.round_timeout,
        )

        try:
            # 步骤 1: 选择 Agent
            selected = await self.select_agents_for_round(round_id)
            training_round.selected_agents = selected

            # 步骤 2: 分发模型
            delivery = await self.distribute_global_model(selected, self.global_model)
            successful_delivery = [
                addr for addr, ok in delivery.items() if ok
            ]
            if len(successful_delivery) < self.min_agents:
                raise RuntimeError(
                    f"模型分发失败: {len(successful_delivery)} < {self.min_agents}"
                )

            # 步骤 3: 收集梯度
            gradients = await self.collect_gradients(
                successful_delivery, timeout=self.round_timeout
            )
            if len(gradients) < self.min_agents:
                raise RuntimeError(
                    f"梯度收集不足: {len(gradients)} < {self.min_agents}"
                )

            training_round.gradients = gradients

            # 步骤 4: 安全聚合
            self.global_model = self.secure_aggregate(
                gradients,
                self.global_model,
                weights=training_round.aggregation_weights,
            )

            # 步骤 5: 记录指标
            stats = self._compute_round_stats(gradients)
            training_round.metrics = stats
            self.training_history.append(training_round)
            self.active_round = training_round

            # 步骤 6: 链上提交聚合结果
            await self._submit_aggregation_to_chain(round_id)

            # 步骤 7: 分发奖励
            await self._distribute_rewards(round_id, gradients.keys())

            # 重置 Agent 状态
            for addr in selected:
                if addr in self.agents:
                    self.agents[addr].status = AgentStatus.IDLE

            logger.info(
                f"第 {round_id} 轮完成,全局模型 v{self.global_model.version}"
            )
            return training_round

        except Exception as e:
            logger.error(f"第 {round_id} 轮训练失败: {e}")
            self.active_round = None
            raise

    async def _submit_aggregation_to_chain(self, round_id: int):
        """将聚合结果提交到 MSG Chain"""
        await self.client.send_tx(
            contract="msg1gradientagg...",
            action="submit_aggregation",
            params={
                "round_id": round_id,
                "model_version": self.global_model.version,
                "model_commitment": self.global_model.commitment(),
                "num_participants": len(self.active_round.gradients if self.active_round else {}),
            },
        )

    async def _distribute_rewards(
        self,
        round_id: int,
        participants: List[str],
    ):
        """向参与训练的 Agent 分发奖励

        奖励计算基于:
        - 梯度大小(反映数据量)
        - 训练时间(反映计算贡献)
        - Agent 信誉分数
        """
        total_reward = 100 * len(participants)  # 基础奖励池
        scores = {}

        for addr in participants:
            info = self.agents[addr]
            gradient_size = len(
                self.active_round.gradients.get(addr, b"")
            ) if self.active_round else 0
            time_factor = max(0.5, 1.0 - info.training_time_ms / 60000)
            score = (
                gradient_size * 0.3
                + time_factor * 0.3
                + info.reputation * 0.4
            )
            scores[addr] = score

        total_score = sum(scores.values())
        if total_score == 0:
            return

        for addr, score in scores.items():
            reward = int(total_reward * (score / total_score))
            await self.client.send_tx(
                contract="msg1flrewards...",
                action="distribute_reward",
                params={
                    "agent": addr,
                    "round_id": round_id,
                    "amount": reward,
                },
            )
            self.agents[addr].contribution_score += reward

    def _compute_round_stats(
        self,
        gradients: Dict[str, bytes],
    ) -> Dict[str, ModelMetrics]:
        """计算本轮训练统计"""
        stats = {}
        for addr, grad_bytes in gradients.items():
            metrics = ModelMetrics()
            metrics.num_samples = len(grad_bytes) // 4
            metrics.round_time_ms = self.agents[addr].training_time_ms
            stats[addr] = metrics
        return stats

    def get_agent_leaderboard(self, top_k: int = 10) -> List[Dict]:
        """获取 Agent 排行榜"""
        sorted_agents = sorted(
            self.agents.values(),
            key=lambda x: x.contribution_score,
            reverse=True,
        )
        return [
            {
                "address": a.address,
                "score": a.contribution_score,
                "reputation": a.reputation,
                "stake": a.stake,
                "rounds_participated": a.last_round,
            }
            for a in sorted_agents[:top_k]
        ]


class DecentralizedFLNetwork:
    """去中心化联邦学习(Gossip 协议)

    在去中心化模式下,没有中心协调器。
    Agent 之间通过 P2P 网络交换模型更新。
    """

    def __init__(self, agent_address: str, p2p_client):
        self.address = agent_address
        self.p2p = p2p_client
        self.local_model = Model()
        self.peers: Dict[str, float] = {}  # address -> trust score
        self.received_models: Dict[int, List[bytes]] = {}
        self.gossip_fanout: int = 3

    async def broadcast_model(self, model: Model):
        """向邻居广播本地模型"""
        message = {
            "type": "model_update",
            "from": self.address,
            "model_weights": model.serialize().hex(),
            "model_commitment": model.commitment(),
            "version": model.version,
            "timestamp": int(time.time()),
        }
        neighbors = self._select_neighbors()
        tasks = []
        for peer in neighbors:
            tasks.append(self.p2p.send(peer, message))
        await asyncio.gather(*tasks, return_exceptions=True)

    def _select_neighbors(self) -> List[str]:
        """基于信誉分数选择邻居节点"""
        sorted_peers = sorted(
            self.peers.items(),
            key=lambda x: x[1],
            reverse=True,
        )
        return [p[0] for p in sorted_peers[:self.gossip_fanout]]

    async def receive_model_update(self, from_addr: str, model_bytes: bytes):
        """接收并验证来自其他 Agent 的模型更新"""
        trust = self.peers.get(from_addr, 0.0)
        if trust < 0.1:
            logger.warning(f"来自 {from_addr} 的更新被拒绝(信誉不足)")
            return

        current_round = int(time.time()) // 3600
        if current_round not in self.received_models:
            self.received_models[current_round] = []
        self.received_models[current_round].append(model_bytes)

        if len(self.received_models[current_round]) >= self.gossip_fanout:
            await self._aggregate_updates(current_round)

    async def _aggregate_updates(self, round_id: int):
        """使用 Krum 算法聚合去中心化更新"""
        updates = self.received_models[round_id]
        if len(updates) < 2:
            return

        n = len(updates)
        distances = np.zeros((n, n))
        for i in range(n):
            for j in range(i + 1, n):
                dist = self._model_distance(updates[i], updates[j])
                distances[i, j] = dist
                distances[j, i] = dist

        score_sum = np.sum(distances, axis=1)
        best_idx = np.argmin(score_sum)

        self.local_model.deserialize(
            updates[best_idx],
            shapes=[(w.shape,) for w in self.local_model.weights],
        )
        self.local_model.version += 1

        del self.received_models[round_id]

    def _model_distance(self, a_bytes: bytes, b_bytes: bytes) -> float:
        """计算两个模型之间的欧几里得距离"""
        a_arr = np.frombuffer(a_bytes, dtype=np.float32)
        b_arr = np.frombuffer(b_bytes, dtype=np.float32)
        return float(np.linalg.norm(a_arr - b_arr))

    async def update_peer_trust(self, peer: str, delta: float):
        """更新对等节点的信任分数"""
        current = self.peers.get(peer, 0.5)
        self.peers[peer] = max(0.0, min(1.0, current + delta))

2.3 Agent 选择策略比较

策略 优点 缺点 适用场景
随机选择 公平,实现简单 可能选择低质量 Agent 同质 Agent 集群
信誉优先 训练质量高 富者愈富,新人难以参与 成熟网络
质押加权 经济安全 资本门槛 DeFi 场景
多样性优先 模型泛化好 收敛慢 异构数据分布
基于贡献 激励参与 计算开销大 成熟生态

2.4 安全聚合(SecAgg)协议

安全聚合确保服务端无法查看单个 Agent 的梯度,只能看到聚合结果:

  1. 密钥协商:Agent 两两之间通过 Diffie-Hellman 协商共享密钥。
  2. 掩码生成:每个 Agent 使用共享密钥生成掩码,对自己梯度加掩。
  3. 聚合:服务端收集所有加掩梯度后,由于掩码抵消,得到原始聚合结果。
import hashlib
from typing import Dict, List


class SecureAggregation:
    """安全聚合(SecAgg)实现

    基于 Bonawitz et al. \"Practical Secure Aggregation for Privacy-Preserving
    Machine Learning\" (CCS 2017) 的掩码方案简化版。
    """

    def __init__(self, num_agents: int, key_size: int = 256):
        self.num_agents = num_agents
        self.key_size = key_size
        self.shared_keys: Dict[tuple, bytes] = {}

    def setup_key_agreement(self, agent_id: int, peer_ids: List[int]):
        """在每个 Agent 上执行密钥协商"""
        for peer_id in peer_ids:
            if peer_id == agent_id:
                continue
            key_material = f"{min(agent_id, peer_id)}:{max(agent_id, peer_id)}"
            shared_key = hashlib.sha256(key_material.encode()).digest()
            self.shared_keys[(agent_id, peer_id)] = shared_key

    def compute_mask(self, agent_id: int, total_agents: int) -> bytes:
        """计算 Agent 的掩码

        对于每个对等节点 j:
        - 如果 j > i: 生成掩码 s_{i,j}
        - 如果 j < i: 生成掩码 s_{j,i} 并在聚合时抵消

        最终掩码 = sum(s_{i,j} for j > i) - sum(s_{j,i} for j < i)
        """
        mask = 0
        for peer_id in range(total_agents):
            if peer_id == agent_id:
                continue
            key = self.shared_keys.get((agent_id, peer_id))
            if key is None:
                continue
            prf = hashlib.shake_256(key).digest(32)
            mask_value = int.from_bytes(prf, "big")

            if peer_id > agent_id:
                mask += mask_value
            else:
                mask -= mask_value

        return mask.to_bytes(32, "big")

    def mask_gradient(
        self,
        gradient_bytes: bytes,
        mask_bytes: bytes,
    ) -> bytes:
        """对梯度应用掩码"""
        grad_int = int.from_bytes(gradient_bytes, "big")
        mask_int = int.from_bytes(mask_bytes, "big")
        masked = (grad_int + mask_int) % (2 ** (len(gradient_bytes) * 8))
        return masked.to_bytes(len(gradient_bytes), "big")

    @staticmethod
    def aggregate_masked_gradients(
        masked_gradients: List[bytes],
        num_agents: int,
    ) -> bytes:
        """聚合所有加掩梯度

        由于掩码互相抵消,聚合结果等于原始梯度和。
        """
        result = 0
        for grad_bytes in masked_gradients:
            result += int.from_bytes(grad_bytes, "big")
        result %= 2 ** (len(masked_gradients[0]) * 8)
        return result.to_bytes(len(masked_gradients[0]), "big")

2.5 通信协议

Agent 与协调器之间的消息格式:

from dataclasses import dataclass
from typing import Optional


@dataclass
class FLMessage:
    """联邦学习消息协议"""
    version: int = 1
    message_id: str = ""
    sender: str = ""
    recipient: str = ""
    round_id: int = 0
    message_type: str = ""
    payload: bytes = b""
    signature: Optional[bytes] = None
    timestamp: int = 0
    nonce: str = ""

    def verify(self) -> bool:
        """验证消息签名"""
        if self.signature is None:
            return False
        from msg_client.crypto import verify_signature
        message_hash = hashlib.sha256(
            f"{self.round_id}:{self.sender}:{self.payload.hex()}:{self.nonce}".encode()
        ).digest()
        return verify_signature(self.sender, message_hash, self.signature)

    def sign(self, private_key: bytes):
        """使用发送方私钥签名"""
        from msg_client.crypto import sign
        message_hash = hashlib.sha256(
            f"{self.round_id}:{self.sender}:{self.payload.hex()}:{self.nonce}".encode()
        ).digest()
        self.signature = sign(private_key, message_hash)

3. 梯度聚合合约

3.1 合约架构

梯度聚合合约是 MSG Chain 上联邦学习的核心链上组件。它负责:

  1. 梯度提交:Agent 将训练后的梯度以交易形式提交到链上。
  2. 贡献验证:验证梯度的完整性和时效性。
  3. 聚合计算:链上执行梯度聚合(或接收链下聚合结果)。
  4. 奖励分发:基于贡献分配代币奖励。

以下使用 CosmWasm 风格的 Rust 实现:

use cosmwasm_std::{
    attr, to_binary, Addr, BankMsg, Binary, Coin, CosmosMsg,
    Decimal, Deps, DepsMut, Env, MessageInfo, Order,
    QueryRequest, Response, StdError, StdResult, Storage,
    Uint128, WasmMsg,
};
use cw_storage_plus::{Item, Map, SnapshotMap};
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use std::collections::HashMap;

// ============================================================
// 合约状态 & 数据结构
// ============================================================

#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, JsonSchema)]
pub struct Config {
    /// 合约管理员
    pub admin: Addr,
    /// 联邦学习协调器地址
    pub coordinator: Addr,
    /// 参与一轮训练的最低 Agent 数量
    pub min_participants: u32,
    /// 一轮训练的超时时间(区块数)
    pub round_timeout: u64,
    /// 基础奖励金额(uMSG)
    pub base_reward: Uint128,
    /// 是否启用差分隐私强制
    pub enforce_differential_privacy: bool,
    /// 最小隐私预算 epsilon
    pub min_epsilon: f64,
    /// 梯度提交的 Gas 上限
    pub max_gradient_size: u64,
}

#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, JsonSchema)]
pub struct RoundInfo {
    /// 轮次编号
    pub round_id: u64,
    /// 开始区块高度
    pub start_height: u64,
    /// 截止区块高度
    pub deadline_height: u64,
    /// 全局模型版本
    pub model_version: u32,
    /// 全局模型承诺(Merkle Root hex)
    pub model_commitment: String,
    /// 参与 Agent 列表
    pub participants: Vec<Addr>,
    /// 已提交 Agent 数量
    pub submitted_count: u32,
    /// 轮次状态
    pub status: RoundStatus,
    /// 聚合结果模型承诺
    pub aggregated_commitment: Option<String>,
}

#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, JsonSchema)]
pub enum RoundStatus {
    Pending,
    Active,
    Aggregated,
    Failed,
    Rewarded,
}

#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, JsonSchema)]
pub struct GradientSubmission {
    /// 提交的 Agent 地址
    pub submitter: Addr,
    /// 轮次编号
    pub round_id: u64,
    /// 梯度数据的哈希(用于验证完整性)
    pub gradient_hash: String,
    /// 梯度数据大小
    pub gradient_size: u64,
    /// 训练样本数量
    pub num_samples: u64,
    /// 训练时间(毫秒)
    pub training_time_ms: u64,
    /// 提交区块高度
    pub submit_height: u64,
    /// 签署证明
    pub attestation: Option<AttestationProof>,
}

#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, JsonSchema)]
pub struct AttestationProof {
    /// TEE 类型
    pub tee_type: String,
    /// 远程证明引用
    pub attestation_ref: String,
    /// 证明数据
    pub quote: Binary,
}

#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, JsonSchema)]
pub struct AgentProfile {
    /// Agent 地址
    pub address: Addr,
    /// 质押金额
    pub stake: Uint128,
    /// 信誉分数 (0.0 - 1.0)
    pub reputation: Decimal,
    /// 累计贡献分数
    pub total_contribution: Uint128,
    /// 参与轮次数量
    pub rounds_participated: u64,
    /// 上次活跃轮次
    pub last_active_round: u64,
    /// 是否被冻结
    pub frozen: bool,
}

// ============================================================
// 存储 (Storage)
// ============================================================

pub const CONFIG: Item<Config> = Item::new("config");
pub const CURRENT_ROUND: Item<u64> = Item::new("current_round");
pub const ROUNDS: Map<u64, RoundInfo> = Map::new("rounds");
pub const AGENTS: Map<&Addr, AgentProfile> = Map::new("agents");
pub const GRADIENT_SUBMISSIONS: Map<(u64, &Addr), GradientSubmission> = Map::new("gradients");
pub const AGENT_LIST: Item<Vec<Addr>> = Item::new("agent_list");

// ============================================================
// 消息类型
// ============================================================

#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, JsonSchema)]
pub struct InstantiateMsg {
    pub admin: String,
    pub coordinator: String,
    pub min_participants: u32,
    pub round_timeout: u64,
    pub base_reward: Uint128,
}

#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, JsonSchema)]
pub enum ExecuteMsg {
    RegisterAgent { agent_address: String, stake: Uint128 },
    UnregisterAgent { agent_address: String },
    StartRound { model_commitment: String },
    SubmitGradient {
        round_id: u64,
        gradient_hash: String,
        gradient_size: u64,
        num_samples: u64,
        training_time_ms: u64,
        attestation: Option<AttestationProof>,
    },
    SubmitAggregation { round_id: u64, aggregated_commitment: String },
    DistributeRewards { round_id: u64 },
    SlashAgent { agent_address: String, reason: String },
    UpdateConfig {
        min_participants: Option<u32>,
        round_timeout: Option<u64>,
        base_reward: Option<Uint128>,
        min_epsilon: Option<f64>,
    },
}

#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, JsonSchema)]
pub enum QueryMsg {
    GetConfig {},
    GetAgent { address: String },
    GetRound { round_id: u64 },
    GetGradientSubmission { round_id: u64, agent: String },
    GetActiveRound {},
    ListAgents { start_after: Option<String>, limit: Option<u32> },
    GetLeaderboard { top_k: Option<u32> },
}

// ============================================================
// 实例化
// ============================================================

pub fn instantiate(
    deps: DepsMut,
    _env: Env,
    _info: MessageInfo,
    msg: InstantiateMsg,
) -> StdResult<Response> {
    let config = Config {
        admin: deps.api.addr_validate(&msg.admin)?,
        coordinator: deps.api.addr_validate(&msg.coordinator)?,
        min_participants: msg.min_participants,
        round_timeout: msg.round_timeout,
        base_reward: msg.base_reward,
        enforce_differential_privacy: true,
        min_epsilon: 1.0,
        max_gradient_size: 1_048_576,
    };

    CONFIG.save(deps.storage, &config)?;
    CURRENT_ROUND.save(deps.storage, &0u64)?;
    AGENT_LIST.save(deps.storage, &vec![])?;

    Ok(Response::new()
        .add_attribute("method", "instantiate")
        .add_attribute("admin", msg.admin)
        .add_attribute("coordinator", msg.coordinator))
}

// ============================================================
// 执行入口
// ============================================================

pub fn execute(
    deps: DepsMut,
    env: Env,
    info: MessageInfo,
    msg: ExecuteMsg,
) -> StdResult<Response> {
    match msg {
        ExecuteMsg::RegisterAgent { agent_address, stake } => {
            execute_register_agent(deps, env, info, agent_address, stake)
        }
        ExecuteMsg::UnregisterAgent { agent_address } => {
            execute_unregister_agent(deps, env, info, agent_address)
        }
        ExecuteMsg::StartRound { model_commitment } => {
            execute_start_round(deps, env, info, model_commitment)
        }
        ExecuteMsg::SubmitGradient { round_id, gradient_hash, gradient_size, num_samples, training_time_ms, attestation } => {
            execute_submit_gradient(deps, env, info, round_id, gradient_hash, gradient_size, num_samples, training_time_ms, attestation)
        }
        ExecuteMsg::SubmitAggregation { round_id, aggregated_commitment } => {
            execute_submit_aggregation(deps, env, info, round_id, aggregated_commitment)
        }
        ExecuteMsg::DistributeRewards { round_id } => {
            execute_distribute_rewards(deps, env, info, round_id)
        }
        ExecuteMsg::SlashAgent { agent_address, reason } => {
            execute_slash_agent(deps, env, info, agent_address, reason)
        }
        ExecuteMsg::UpdateConfig { min_participants, round_timeout, base_reward, min_epsilon } => {
            execute_update_config(deps, env, info, min_participants, round_timeout, base_reward, min_epsilon)
        }
    }
}

// ============================================================
// 注册 Agent
// ============================================================

pub fn execute_register_agent(
    deps: DepsMut,
    _env: Env,
    info: MessageInfo,
    agent_address: String,
    stake: Uint128,
) -> StdResult<Response> {
    let config = CONFIG.load(deps.storage)?;
    if info.sender != config.admin && info.sender.as_str() != agent_address {
        return Err(StdError::generic_err("未授权:仅管理员或本人可注册"));
    }

    if stake < Uint128::from(100u128) {
        return Err(StdError::generic_err("质押不足:至少需要 100 uMSG"));
    }

    let addr = deps.api.addr_validate(&agent_address)?;

    if AGENTS.has(deps.storage, &addr) {
        return Err(StdError::generic_err("Agent 已注册"));
    }

    let agent = AgentProfile {
        address: addr.clone(),
        stake,
        reputation: Decimal::percent(50),
        total_contribution: Uint128::zero(),
        rounds_participated: 0,
        last_active_round: 0,
        frozen: false,
    };

    AGENTS.save(deps.storage, &addr, &agent)?;

    let mut list = AGENT_LIST.load(deps.storage)?;
    list.push(addr);
    AGENT_LIST.save(deps.storage, &list)?;

    Ok(Response::new()
        .add_attribute("method", "register_agent")
        .add_attribute("agent", agent_address)
        .add_attribute("stake", stake.to_string()))
}

// ============================================================
// 注销 Agent
// ============================================================

pub fn execute_unregister_agent(
    deps: DepsMut,
    _env: Env,
    info: MessageInfo,
    agent_address: String,
) -> StdResult<Response> {
    let config = CONFIG.load(deps.storage)?;
    if info.sender != config.admin && info.sender.as_str() != agent_address {
        return Err(StdError::generic_err("未授权"));
    }

    let addr = deps.api.addr_validate(&agent_address)?;
    let agent = AGENTS.load(deps.storage, &addr)?;

    let refund = BankMsg::Send {
        to_address: agent_address.clone(),
        amount: vec![Coin {
            denom: "uMSG".to_string(),
            amount: agent.stake,
        }],
    };

    AGENTS.remove(deps.storage, &addr);

    let mut list = AGENT_LIST.load(deps.storage)?;
    list.retain(|a| a.as_str() != agent_address);
    AGENT_LIST.save(deps.storage, &list)?;

    Ok(Response::new()
        .add_message(refund)
        .add_attribute("method", "unregister_agent")
        .add_attribute("agent", agent_address))
}

// ============================================================
// 开始新一轮训练
// ============================================================

pub fn execute_start_round(
    deps: DepsMut,
    env: Env,
    info: MessageInfo,
    model_commitment: String,
) -> StdResult<Response> {
    let config = CONFIG.load(deps.storage)?;
    if info.sender != config.coordinator {
        return Err(StdError::generic_err("仅协调器可启动轮次"));
    }

    let current_round = CURRENT_ROUND.load(deps.storage)?;
    let new_round_id = current_round + 1;

    let agent_list = AGENT_LIST.load(deps.storage)?;
    let mut participants = Vec::new();
    for addr in &agent_list {
        let agent = AGENTS.load(deps.storage, addr)?;
        if !agent.frozen && agent.stake > Uint128::zero() {
            participants.push(addr.clone());
        }
    }

    if (participants.len() as u32) < config.min_participants {
        return Err(StdError::generic_err(format!(
            "活跃 Agent 不足: {} < {}",
            participants.len(),
            config.min_participants
        )));
    }

    let round = RoundInfo {
        round_id: new_round_id,
        start_height: env.block.height,
        deadline_height: env.block.height + config.round_timeout,
        model_version: new_round_id as u32,
        model_commitment: model_commitment.clone(),
        participants,
        submitted_count: 0,
        status: RoundStatus::Active,
        aggregated_commitment: None,
    };

    ROUNDS.save(deps.storage, new_round_id, &round)?;
    CURRENT_ROUND.save(deps.storage, &new_round_id)?;

    Ok(Response::new()
        .add_attribute("method", "start_round")
        .add_attribute("round_id", new_round_id.to_string())
        .add_attribute("num_participants", round.participants.len().to_string()))
}

// ============================================================
// 提交梯度
// ============================================================

pub fn execute_submit_gradient(
    deps: DepsMut,
    env: Env,
    info: MessageInfo,
    round_id: u64,
    gradient_hash: String,
    gradient_size: u64,
    num_samples: u64,
    training_time_ms: u64,
    attestation: Option<AttestationProof>,
) -> StdResult<Response> {
    let config = CONFIG.load(deps.storage)?;

    let mut round = ROUNDS.load(deps.storage, round_id)?;
    if round.status != RoundStatus::Active {
        return Err(StdError::generic_err("本轮已结束或不存在"));
    }

    if env.block.height > round.deadline_height {
        round.status = RoundStatus::Failed;
        ROUNDS.save(deps.storage, round_id, &round)?;
        return Err(StdError::generic_err("本轮已超时"));
    }

    if !round.participants.contains(&info.sender) {
        return Err(StdError::generic_err("提交者不在参与者列表中"));
    }

    let submission_key = (round_id, &info.sender);
    if GRADIENT_SUBMISSIONS.has(deps.storage, submission_key) {
        return Err(StdError::generic_err("已提交过梯度,不可重复提交"));
    }

    if gradient_size > config.max_gradient_size {
        return Err(StdError::generic_err("梯度数据大小超限"));
    }

    if config.enforce_differential_privacy {
        if let Some(ref attest) = attestation {
            verify_attestation(deps, &config, &info.sender, attest)?;
        }
    }

    let submission = GradientSubmission {
        submitter: info.sender.clone(),
        round_id,
        gradient_hash: gradient_hash.clone(),
        gradient_size,
        num_samples,
        training_time_ms,
        submit_height: env.block.height,
        attestation,
    };

    GRADIENT_SUBMISSIONS.save(deps.storage, submission_key, &submission)?;

    round.submitted_count += 1;
    ROUNDS.save(deps.storage, round_id, &round)?;

    AGENTS.update(deps.storage, &info.sender, |opt| -> StdResult<AgentProfile> {
        let mut agent = opt.ok_or_else(|| StdError::generic_err("Agent 不存在"))?;
        agent.rounds_participated += 1;
        agent.last_active_round = round_id;
        Ok(agent)
    })?;

    Ok(Response::new()
        .add_attribute("method", "submit_gradient")
        .add_attribute("round_id", round_id.to_string())
        .add_attribute("submitter", info.sender.as_str())
        .add_attribute("gradient_hash", &gradient_hash)
        .add_attribute("gradient_size", gradient_size.to_string()))
}

// ============================================================
// 提交聚合结果
// ============================================================

pub fn execute_submit_aggregation(
    deps: DepsMut,
    env: Env,
    info: MessageInfo,
    round_id: u64,
    aggregated_commitment: String,
) -> StdResult<Response> {
    let config = CONFIG.load(deps.storage)?;
    if info.sender != config.coordinator {
        return Err(StdError::generic_err("仅协调器可提交聚合结果"));
    }

    let mut round = ROUNDS.load(deps.storage, round_id)?;
    if round.status != RoundStatus::Active {
        return Err(StdError::generic_err("轮次状态不正确"));
    }

    if round.submitted_count < config.min_participants {
        round.status = RoundStatus::Failed;
        ROUNDS.save(deps.storage, round_id, &round)?;
        return Err(StdError::generic_err(format!(
            "提交数不足: {} < {}",
            round.submitted_count, config.min_participants
        )));
    }

    round.status = RoundStatus::Aggregated;
    round.aggregated_commitment = Some(aggregated_commitment.clone());
    ROUNDS.save(deps.storage, round_id, &round)?;

    Ok(Response::new()
        .add_attribute("method", "submit_aggregation")
        .add_attribute("round_id", round_id.to_string())
        .add_attribute("aggregated_commitment", aggregated_commitment))
}

// ============================================================
// 分发奖励
// ============================================================

pub fn execute_distribute_rewards(
    deps: DepsMut,
    env: Env,
    info: MessageInfo,
    round_id: u64,
) -> StdResult<Response> {
    let config = CONFIG.load(deps.storage)?;
    if info.sender != config.coordinator && info.sender != config.admin {
        return Err(StdError::generic_err("未授权"));
    }

    let mut round = ROUNDS.load(deps.storage, round_id)?;
    if round.status != RoundStatus::Aggregated {
        return Err(StdError::generic_err("轮次尚未聚合完成"));
    }

    let mut contributions: Vec<(Addr, Uint128, f64)> = Vec::new();
    let mut total_score = 0f64;

    for participant in &round.participants {
        let sub_key = (round_id, participant);
        if let Ok(submission) = GRADIENT_SUBMISSIONS.load(deps.storage, sub_key) {
            let agent = AGENTS.load(deps.storage, participant)?;
            let time_factor = 1.0 / (1.0 + (submission.training_time_ms as f64 / 60000.0));
            let sample_factor = (submission.num_samples as f64).ln_1p();
            let reputation_factor = agent.reputation.to_f64();
            let score = sample_factor * reputation_factor * time_factor;

            total_score += score;
            contributions.push((participant.clone(), submission.num_samples, score));
        }
    }

    if total_score <= 0.0 {
        return Err(StdError::generic_err("无有效贡献可奖励"));
    }

    let mut messages: Vec<CosmosMsg> = Vec::new();
    let total_reward = config.base_reward * Uint128::from(round.participants.len() as u128);

    for (addr, _samples, score) in &contributions {
        let reward_amount = total_reward * Uint128::from((*score / total_score * 1e6) as u128)
            / Uint128::from(1_000_000u128);

        if reward_amount > Uint128::zero() {
            messages.push(CosmosMsg::Bank(BankMsg::Send {
                to_address: addr.to_string(),
                amount: vec![Coin {
                    denom: "uMSG".to_string(),
                    amount: reward_amount,
                }],
            }));

            AGENTS.update(deps.storage, addr, |opt| -> StdResult<AgentProfile> {
                let mut agent = opt.ok_or_else(|| StdError::generic_err("Agent 不存在"))?;
                agent.total_contribution += reward_amount;
                let rep_increment = Decimal::from_ratio(1u128, 1000u128);
                agent.reputation = Decimal::from_atomics(
                    agent.reputation.atomics() + rep_increment.atomics(),
                    agent.reputation.decimal_places(),
                ).unwrap_or(agent.reputation);
                if agent.reputation > Decimal::percent(100) {
                    agent.reputation = Decimal::percent(100);
                }
                Ok(agent)
            })?;
        }
    }

    round.status = RoundStatus::Rewarded;
    ROUNDS.save(deps.storage, round_id, &round)?;

    Ok(Response::new()
        .add_messages(messages)
        .add_attribute("method", "distribute_rewards")
        .add_attribute("round_id", round_id.to_string())
        .add_attribute("num_rewarded", contributions.len().to_string()))
}

// ============================================================
// Slash 惩罚
// ============================================================

pub fn execute_slash_agent(
    deps: DepsMut,
    _env: Env,
    info: MessageInfo,
    agent_address: String,
    reason: String,
) -> StdResult<Response> {
    let config = CONFIG.load(deps.storage)?;
    if info.sender != config.admin && info.sender != config.coordinator {
        return Err(StdError::generic_err("未授权"));
    }

    let addr = deps.api.addr_validate(&agent_address)?;
    let mut agent = AGENTS.load(deps.storage, &addr)?;

    let slash_amount = agent.stake / Uint128::from(10u128);
    agent.stake = agent.stake.checked_sub(slash_amount)?;
    agent.reputation = Decimal::from_atomics(
        agent.reputation.atomics().saturating_sub(
            Decimal::percent(10).atomics(),
        ),
        agent.reputation.decimal_places(),
    ).unwrap_or(Decimal::zero());

    if agent.stake.is_zero() || agent.reputation.is_zero() {
        agent.frozen = true;
    }

    AGENTS.save(deps.storage, &addr, &agent)?;

    Ok(Response::new()
        .add_attribute("method", "slash_agent")
        .add_attribute("agent", agent_address)
        .add_attribute("slash_amount", slash_amount.to_string())
        .add_attribute("reason", reason))
}

// ============================================================
// 更新配置
// ============================================================

pub fn execute_update_config(
    deps: DepsMut,
    _env: Env,
    info: MessageInfo,
    min_participants: Option<u32>,
    round_timeout: Option<u64>,
    base_reward: Option<Uint128>,
    min_epsilon: Option<f64>,
) -> StdResult<Response> {
    let mut config = CONFIG.load(deps.storage)?;
    if info.sender != config.admin {
        return Err(StdError::generic_err("仅管理员可更新配置"));
    }

    if let Some(val) = min_participants { config.min_participants = val; }
    if let Some(val) = round_timeout { config.round_timeout = val; }
    if let Some(val) = base_reward { config.base_reward = val; }
    if let Some(val) = min_epsilon { config.min_epsilon = val; }

    CONFIG.save(deps.storage, &config)?;
    Ok(Response::new().add_attribute("method", "update_config"))
}

// ============================================================
// 查询
// ============================================================

pub fn query(deps: Deps, _env: Env, msg: QueryMsg) -> StdResult<Binary> {
    match msg {
        QueryMsg::GetConfig {} => to_binary(&query_config(deps)?),
        QueryMsg::GetAgent { address } => to_binary(&query_agent(deps, address)?),
        QueryMsg::GetRound { round_id } => to_binary(&query_round(deps, round_id)?),
        QueryMsg::GetGradientSubmission { round_id, agent } => {
            to_binary(&query_gradient_submission(deps, round_id, agent)?)
        }
        QueryMsg::GetActiveRound {} => to_binary(&query_active_round(deps)?),
        QueryMsg::ListAgents { start_after, limit } => {
            to_binary(&query_list_agents(deps, start_after, limit)?)
        }
        QueryMsg::GetLeaderboard { top_k } => {
            to_binary(&query_leaderboard(deps, top_k)?)
        }
    }
}

pub fn query_config(deps: Deps) -> StdResult<Config> { CONFIG.load(deps.storage) }

pub fn query_agent(deps: Deps, address: String) -> StdResult<AgentProfile> {
    let addr = deps.api.addr_validate(&address)?;
    AGENTS.load(deps.storage, &addr)
}

pub fn query_round(deps: Deps, round_id: u64) -> StdResult<RoundInfo> {
    ROUNDS.load(deps.storage, round_id)
}

pub fn query_gradient_submission(deps: Deps, round_id: u64, agent: String) -> StdResult<GradientSubmission> {
    let addr = deps.api.addr_validate(&agent)?;
    GRADIENT_SUBMISSIONS.load(deps.storage, (round_id, &addr))
}

pub fn query_active_round(deps: Deps) -> StdResult<Option<RoundInfo>> {
    let current_round = CURRENT_ROUND.load(deps.storage)?;
    if current_round == 0 { return Ok(None); }
    let round = ROUNDS.load(deps.storage, current_round)?;
    if round.status == RoundStatus::Active { Ok(Some(round)) } else { Ok(None) }
}

pub fn query_list_agents(deps: Deps, start_after: Option<String>, limit: Option<u32>) -> StdResult<Vec<AgentProfile>> {
    let limit = limit.unwrap_or(30).min(100) as usize;
    let start = start_after.map(|s| deps.api.addr_validate(&s)).transpose()?;
    let agents: StdResult<Vec<AgentProfile>> = AGENTS
        .range(deps.storage, start.as_ref(), None, Order::Ascending)
        .take(limit)
        .map(|item| item.map(|(_, agent)| agent))
        .collect();
    agents
}

pub fn query_leaderboard(deps: Deps, top_k: Option<u32>) -> StdResult<Vec<AgentProfile>> {
    let top_k = top_k.unwrap_or(10) as usize;
    let mut agents: Vec<AgentProfile> = AGENTS
        .range(deps.storage, None, None, Order::Ascending)
        .filter_map(|item| item.ok())
        .map(|(_, agent)| agent)
        .collect();
    agents.sort_by(|a, b| b.total_contribution.cmp(&a.total_contribution));
    agents.truncate(top_k);
    Ok(agents)
}

// ============================================================
// TEE 远程证明验证
// ============================================================

fn verify_attestation(deps: Deps, config: &Config, agent: &Addr, attestation: &AttestationProof) -> StdResult<()> {
    match attestation.tee_type.as_str() {
        "sgx" | "tdx" | "nitro" => {}
        _ => return Err(StdError::generic_err("不支持的 TEE 类型")),
    }
    if attestation.quote.is_empty() {
        return Err(StdError::generic_err("远程证明数据为空"));
    }
    Ok(())
}

// ============================================================
// 测试
// ============================================================

#[cfg(test)]
mod tests {
    use super::*;
    use cosmwasm_std::testing::{mock_dependencies, mock_env, mock_info};
    use cosmwasm_std::{coins, from_binary};

    fn init_contract(deps: DepsMut) -> Response {
        let msg = InstantiateMsg {
            admin: "msg1admin...".to_string(),
            coordinator: "msg1coordinator...".to_string(),
            min_participants: 3,
            round_timeout: 100,
            base_reward: Uint128::from(1000u128),
        };
        instantiate(deps, mock_env(), mock_info("msg1admin...", &[]), msg).unwrap()
    }

    #[test]
    fn test_register_agent() {
        let mut deps = mock_dependencies();
        init_contract(deps.as_mut());
        let msg = ExecuteMsg::RegisterAgent {
            agent_address: "msg1agent1...".to_string(),
            stake: Uint128::from(1000u128),
        };
        let res = execute(deps.as_mut(), mock_env(), mock_info("msg1agent1...", &[]), msg).unwrap();
        assert_eq!(res.attributes[0].value, "register_agent");
        assert_eq!(res.attributes[1].value, "msg1agent1...");
    }

    #[test]
    fn test_register_agent_insufficient_stake() {
        let mut deps = mock_dependencies();
        init_contract(deps.as_mut());
        let msg = ExecuteMsg::RegisterAgent {
            agent_address: "msg1agent1...".to_string(),
            stake: Uint128::from(50u128),
        };
        let err = execute(deps.as_mut(), mock_env(), mock_info("msg1agent1...", &[]), msg).unwrap_err();
        assert!(err.to_string().contains("质押不足"));
    }

    #[test]
    fn test_submit_gradient() {
        let mut deps = mock_dependencies();
        init_contract(deps.as_mut());
        for i in 1..=4 {
            execute(deps.as_mut(), mock_env(), mock_info(&format!("msg1agent{}...", i), &[]),
                ExecuteMsg::RegisterAgent { agent_address: format!("msg1agent{}...", i), stake: Uint128::from(1000u128) }).unwrap();
        }
        execute(deps.as_mut(), mock_env(), mock_info("msg1coordinator...", &[]),
            ExecuteMsg::StartRound { model_commitment: "abc123".to_string() }).unwrap();
        let res = execute(deps.as_mut(), mock_env(), mock_info("msg1agent1...", &[]),
            ExecuteMsg::SubmitGradient { round_id: 1, gradient_hash: "grad_hash_1".to_string(), gradient_size: 1024, num_samples: 100, training_time_ms: 5000, attestation: None }).unwrap();
        assert_eq!(res.attributes[0].value, "submit_gradient");
    }

    #[test]
    fn test_leaderboard() {
        let mut deps = mock_dependencies();
        init_contract(deps.as_mut());
        for addr in &["msg1agent_a...", "msg1agent_b...", "msg1agent_c..."] {
            execute(deps.as_mut(), mock_env(), mock_info(addr, &[]),
                ExecuteMsg::RegisterAgent { agent_address: addr.to_string(), stake: Uint128::from(1000u128) }).unwrap();
        }
        let res = query(deps.as_ref(), mock_env(), QueryMsg::GetLeaderboard { top_k: Some(5) }).unwrap();
        let board: Vec<AgentProfile> = from_binary(&res).unwrap();
        assert_eq!(board.len(), 3);
    }
}

3.2 部署指南

# 编译合约
cargo wasm

# 优化体积
docker run --rm -v "$(pwd)":/code \
  --mount type=volume,source="$(basename "$(pwd)")_cache",target=/code/target \
  --mount type=volume,source=registry_cache,target=/usr/local/cargo/registry \
  cosmwasm/workspace-optimizer:0.12.13

# 部署到 MSG Chain
msgd tx wasm store artifacts/gradient_aggregator.wasm \
  --from admin \
  --chain-id msg-chain-1 \
  --gas auto --fees 5000uMSG

# 实例化
msgd tx wasm instantiate CODE_ID \
  '{"admin":"msg1admin...","coordinator":"msg1coordinator...","min_participants":3,"round_timeout":100,"base_reward":"1000"}' \
  --from admin \
  --label "FL Gradient Aggregator" \
  --chain-id msg-chain-1

3.3 事件(Events)

合约在关键操作时发出事件,供监听器消费:

事件名称 属性 触发时机
agent_registered agent, stake Agent 注册成功
agent_unregistered agent Agent 注销
round_started round_id, num_participants 训练轮次开始
gradient_submitted round_id, submitter, gradient_hash 梯度提交成功
aggregation_submitted round_id, aggregated_commitment 聚合结果提交
rewards_distributed round_id, num_rewarded 奖励分发完成
agent_slashed agent, amount, reason Agent 被惩罚

4. 差分隐私保护

4.1 原理

差分隐私(Differential Privacy, DP)通过在梯度中添加精心校准的噪声,使得攻击者无法判断某个特定样本是否参与了训练。形式化定义为:

一个随机化机制 M 满足 (ε, δ)-差分隐私,如果对于任意相邻数据集 D 和 D'(相差一条记录),以及任意输出集合 S:

Pr[M(D) ∈ S] ≤ e^ε · Pr[M(D') ∈ S] + δ

其中:

4.2 完整实现

import numpy as np
import hashlib
from typing import Dict, List, Optional, Tuple, Callable
from dataclasses import dataclass, field
from enum import Enum


class DPMechanism(Enum):
    LAPLACE = "laplace"
    GAUSSIAN = "gaussian"
    EXPONENTIAL = "exponential"


@dataclass
class PrivacyBudget:
    total_epsilon: float = 0.0
    total_delta: float = 0.0
    remaining_budget: float = 1.0
    spent_rounds: List[float] = field(default_factory=list)

    def can_spend(self, epsilon: float, delta: float = 1e-5) -> bool:
        required = self._composition_cost(epsilon)
        return self.remaining_budget >= required

    def spend(self, epsilon: float, delta: float = 1e-5):
        cost = self._composition_cost(epsilon)
        if cost > self.remaining_budget:
            raise ValueError(f"隐私预算不足: 需要 {cost:.4f}, 剩余 {self.remaining_budget:.4f}")
        self.total_epsilon += epsilon
        self.total_delta += delta
        self.remaining_budget -= cost
        self.spent_rounds.append(epsilon)

    def _composition_cost(self, epsilon: float) -> float:
        k = len(self.spent_rounds) + 1
        delta_prime = 1e-6
        return epsilon * np.sqrt(2 * k * np.log(1 / delta_prime))

    def reset(self):
        self.total_epsilon = 0.0
        self.total_delta = 0.0
        self.remaining_budget = 1.0
        self.spent_rounds = []

    def to_dict(self) -> dict:
        return {
            "total_epsilon": self.total_epsilon,
            "total_delta": self.total_delta,
            "remaining_budget": self.remaining_budget,
            "rounds": len(self.spent_rounds),
        }


class DifferentialPrivacy:
    def __init__(
        self,
        mechanism: DPMechanism = DPMechanism.GAUSSIAN,
        budget: Optional[PrivacyBudget] = None,
        clip_norm: float = 1.0,
        secure_rng: bool = True,
    ):
        self.mechanism = mechanism
        self.budget = budget or PrivacyBudget()
        self.clip_norm = clip_norm
        self.secure_rng = secure_rng
        self.rng = np.random.SeedSequence(
            int.from_bytes(hashlib.sha256(b"msg_chain_dp_seed").digest(), "big")
        )
        self._bit_generator = np.random.MT19937(self.rng)

    def apply_dp(
        self,
        gradients: np.ndarray,
        epsilon: float = 1.0,
        delta: float = 1e-5,
    ) -> np.ndarray:
        clipped = self._clip_gradients(gradients, self.clip_norm)

        if self.mechanism == DPMechanism.LAPLACE:
            noised = self._laplace_mechanism(clipped, epsilon)
        elif self.mechanism == DPMechanism.GAUSSIAN:
            noised = self._gaussian_mechanism(clipped, epsilon, delta)
        else:
            raise ValueError(f"不支持的机制: {self.mechanism}")

        self.budget.spend(epsilon, delta)
        return noised

    def _clip_gradients(self, gradients: np.ndarray, clip_norm: float) -> np.ndarray:
        norm = np.linalg.norm(gradients)
        if norm > clip_norm:
            return gradients * (clip_norm / norm)
        return gradients

    def _compute_sensitivity(self, gradients: np.ndarray) -> float:
        return self.clip_norm

    def _laplace_mechanism(self, values: np.ndarray, epsilon: float) -> np.ndarray:
        sensitivity = self._compute_sensitivity(values)
        scale = sensitivity / epsilon
        noise = np.random.default_rng(self._bit_generator).laplace(0, scale, values.shape)
        return values + noise

    def _gaussian_mechanism(self, values: np.ndarray, epsilon: float, delta: float) -> np.ndarray:
        sensitivity = self._compute_sensitivity(values)
        sigma = sensitivity * np.sqrt(2 * np.log(1.25 / delta)) / epsilon
        noise = np.random.default_rng(self._bit_generator).normal(0, sigma, values.shape)
        return values + noise

    def apply_dp_to_model(self, model_weights: List[np.ndarray], epsilon: float, delta: float = 1e-5) -> List[np.ndarray]:
        num_layers = len(model_weights)
        per_layer_epsilon = epsilon / num_layers
        per_layer_delta = delta / num_layers
        noised_weights = []
        for layer_weights in model_weights:
            noised = self.apply_dp(layer_weights, per_layer_epsilon, per_layer_delta)
            noised_weights.append(noised)
        return noised_weights

    def compose_with_rdp(self, epsilons: List[float], deltas: List[float], order: Optional[float] = None) -> Tuple[float, float]:
        if order is None:
            order = self._optimal_rdp_order(epsilons)
        rdp_sum = sum(self._eps_to_rdp(eps, order) for eps in epsilons)
        total_delta = sum(deltas)
        total_epsilon = rdp_sum + np.log(1 / total_delta) / (order - 1)
        return total_epsilon, total_delta

    def _eps_to_rdp(self, epsilon: float, alpha: float) -> float:
        return epsilon * alpha / 2

    def _optimal_rdp_order(self, epsilons: List[float]) -> float:
        candidates = np.linspace(2, 50, 49)
        best_order = 2.0
        best_eps = float('inf')
        for alpha in candidates:
            rdp_sum = sum(self._eps_to_rdp(eps, alpha) for eps in epsilons)
            delta_prime = 1e-6
            total_eps = rdp_sum + np.log(1 / delta_prime) / (alpha - 1)
            if total_eps < best_eps:
                best_eps = total_eps
                best_order = alpha
        return best_order


class AdaptiveDPOptimizer:
    def __init__(self, dp: DifferentialPrivacy, initial_epsilon: float = 1.0, min_epsilon: float = 0.1, max_epsilon: float = 5.0, adaptation_rate: float = 0.5):
        self.dp = dp
        self.current_epsilon = initial_epsilon
        self.min_epsilon = min_epsilon
        self.max_epsilon = max_epsilon
        self.adaptation_rate = adaptation_rate
        self.loss_history: List[float] = []

    def get_epsilon_for_round(self, current_loss: float) -> float:
        self.loss_history.append(current_loss)
        if len(self.loss_history) < 2:
            return self.current_epsilon
        loss_delta = self.loss_history[-2] - self.loss_history[-1]
        relative_change = loss_delta / max(self.loss_history[-2], 1e-8)
        if relative_change > 0.05:
            self.current_epsilon = min(self.max_epsilon, self.current_epsilon * (1 + self.adaptation_rate))
        elif relative_change < -0.01:
            self.current_epsilon = max(self.min_epsilon, self.current_epsilon * (1 - self.adaptation_rate))
        return self.current_epsilon


class LocalDPAgent:
    def __init__(self, agent_id: str, dp_config: Optional[DifferentialPrivacy] = None, epsilon_per_round: float = 0.1):
        self.agent_id = agent_id
        self.dp = dp_config or DifferentialPrivacy(mechanism=DPMechanism.GAUSSIAN, clip_norm=1.0)
        self.epsilon_per_round = epsilon_per_round

    def protect_gradients(self, raw_gradients: List[np.ndarray], batch_size: int) -> List[np.ndarray]:
        protected = []
        for layer_grad in raw_gradients:
            if layer_grad.ndim > 1:
                norms = np.linalg.norm(layer_grad.reshape(layer_grad.shape[0], -1), axis=1)
                multipliers = np.minimum(1.0, self.dp.clip_norm / (norms + 1e-8))
                clipped = layer_grad * multipliers.reshape(-1, *([1] * (layer_grad.ndim - 1)))
            else:
                clipped = self.dp._clip_gradients(layer_grad, self.dp.clip_norm)
            sensitivity = self.dp.clip_norm / batch_size
            sigma = sensitivity * np.sqrt(2 * np.log(1.25 / 1e-5)) / self.epsilon_per_round
            noise = np.random.normal(0, sigma, clipped.shape)
            noisy_grad = clipped / batch_size + noise
            protected.append(noisy_grad)
        return protected


class PrivacyAuditor:
    def __init__(self):
        self.agent_budgets: Dict[str, PrivacyBudget] = {}
        self.global_audit_log: List[dict] = []

    def register_agent(self, agent_id: str, initial_budget: float = 1.0):
        budget = PrivacyBudget()
        budget.remaining_budget = initial_budget
        self.agent_budgets[agent_id] = budget

    def log_round(self, round_id: int, agent_id: str, epsilon_spent: float, num_samples: int):
        if agent_id not in self.agent_budgets:
            raise ValueError(f"Agent {agent_id} 未注册")
        budget = self.agent_budgets[agent_id]
        budget.spend(epsilon_spent)
        self.global_audit_log.append({
            "round_id": round_id,
            "agent_id": agent_id,
            "epsilon_spent": epsilon_spent,
            "remaining_budget": budget.remaining_budget,
            "timestamp": __import__("time").time(),
        })

    def generate_audit_report(self) -> dict:
        total_agents = len(self.agent_budgets)
        avg_remaining = np.mean([b.remaining_budget for b in self.agent_budgets.values()])
        return {
            "total_rounds_logged": len(self.global_audit_log),
            "total_agents": total_agents,
            "average_remaining_budget": avg_remaining,
            "agents_at_risk": sum(1 for b in self.agent_budgets.values() if b.remaining_budget < 0.1),
            "compliance_status": "compliant" if avg_remaining > 0.05 else "at_risk",
            "log": self.global_audit_log[-100:],
        }

    def to_chain_payload(self) -> dict:
        return {
            "total_epsilon": sum(b.total_epsilon for b in self.agent_budgets.values()),
            "num_agents": len(self.agent_budgets),
            "compliance_hash": hashlib.sha256(str(self.generate_audit_report()).encode()).hexdigest(),
        }

4.3 噪声校准指南

不同 ε 值对模型精度和隐私保护的影响:

ε 隐私级别 噪声幅度 典型应用 精度损失
0.01 - 0.1 强保护 大 医疗数据、金融数据 5-15%
0.1 - 1.0 中等保护 中 用户行为分析 2-8%
1.0 - 10.0 轻量保护 小 聚合统计 <2%
>10.0 弱保护 极小 公开数据 可忽略

4.4 隐私预算追踪流程

每轮训练流程:

Agent 本地:
  1. 检查隐私预算剩余
  2. 计算本轮所需 ε
  3. 梯度裁剪 + 加噪
  4. 更新本地预算记录

链上验证:
  5. Agent 提交梯度 + 隐私证明
  6. 合约验证 ε 合规 (≥ min_epsilon)
  7. 记录 Agent 隐私花费

审计:
  8. 定期生成隐私报告
  9. 公开验证合规性

5. TEE 机密计算集成

5.1 为什么需要 TEE

可信执行环境(Trusted Execution Environment, TEE)在联邦学习中扮演关键角色:

  1. 隔离执行:Agent 的训练过程在硬件级隔离的 enclave 中执行,即使主机操作系统被攻破,也无法查看训练数据或模型参数。
  2. 远程证明(Remote Attestation):其他参与者可以验证 Agent 的代码确实在 genuine TEE 中运行且未被篡改。
  3. 机密性:数据在传输和内存中始终加密,只有 enclave 内部可以解密。

MSG Chain 支持的 TEE 类型:

类型 硬件 特点
Intel SGX Intel CPU 应用级 enclave,内存受限(128-512MB EPC)
Intel TDX Intel CPU 虚拟机级 TEE,支持完整 OS
AMD SEV-SNP AMD CPU 虚拟机级 TEE,大内存支持
AWS Nitro Nitro 芯片 云原生 TEE,管理方便

5.2 完整实现

import asyncio
import json
import logging
import os
import struct
import time
from dataclasses import dataclass
from typing import Dict, List, Optional, Any, Callable
from enum import Enum

logger = logging.getLogger(__name__)


class TEEType(Enum):
    SGX = "sgx"
    TDX = "tdx"
    NITRO = "nitro"
    SEV_SNP = "sev_snp"


@dataclass
class AttestationEvidence:
    """远程证明证据"""
    tee_type: TEEType
    quote: bytes
    enclave_hash: str
    signer: str
    mr_signer: str
    mr_enclave: str
    is_debug: bool
    timestamp: int
    user_data: bytes


@dataclass
class AttestationResult:
    verified: bool
    tee_type: TEEType
    is_production: bool
    platform_ok: bool
    code_ok: bool
    details: str


class TEEEnclave:
    def __init__(self, tee_type: TEEType = TEEType.SGX, enclave_path: Optional[str] = None, spid: Optional[str] = None, ias_key: Optional[str] = None):
        self.tee_type = tee_type
        self.enclave_path = enclave_path
        self.spid = spid
        self.ias_key = ias_key
        self._is_initialized = False
        self._measurement: Optional[str] = None
        self._sealed_data: Dict[str, bytes] = {}

    async def initialize(self) -> bool:
        logger.info(f"初始化 {self.tee_type.value} Enclave...")
        await asyncio.sleep(0.5)
        if self.tee_type == TEEType.SGX: self._init_sgx()
        elif self.tee_type == TEEType.TDX: self._init_tdx()
        elif self.tee_type == TEEType.NITRO: self._init_nitro()
        elif self.tee_type == TEEType.SEV_SNP: self._init_sev_snp()
        self._is_initialized = True
        self._measurement = self._compute_measurement()
        logger.info(f"Enclave 初始化完成, 度量值: {self._measurement[:16]}...")
        return True

    def _init_sgx(self): self._measurement = "sgx_measurement"; self._max_epc_size = 128 * 1024 * 1024
    def _init_tdx(self): self._measurement = "tdx_measurement"; self._max_epc_size = 512 * 1024 * 1024
    def _init_nitro(self): self._measurement = "nitro_measurement"; self._max_epc_size = 1024 * 1024 * 1024
    def _init_sev_snp(self): self._measurement = "sev_measurement"; self._max_epc_size = 1024 * 1024 * 1024

    def _compute_measurement(self) -> str:
        import hashlib
        h = hashlib.sha256()
        h.update(self.tee_type.value.encode())
        h.update(b"msg_chain_fl_v1.0")
        h.update(os.urandom(8))
        return h.hexdigest()

    @property
    def is_initialized(self) -> bool: return self._is_initialized
    @property
    def measurement(self) -> Optional[str]: return self._measurement

    async def execute(self, command: str, params: dict) -> bytes:
        if not self._is_initialized:
            raise RuntimeError("Enclave 未初始化")
        logger.info(f"Enclave 执行: {command}")
        if command == "train": return await self._enclave_train(params)
        elif command == "inference": return await self._enclave_inference(params)
        elif command == "encrypt": return self._enclave_encrypt(params)
        elif command == "decrypt": return self._enclave_decrypt(params)
        else: raise ValueError(f"未知命令: {command}")

    async def _enclave_train(self, params: dict) -> bytes:
        model_bytes = params.get("model", b"")
        encrypted_data = params.get("data", b"")
        encryption_key = params.get("key", b"")
        decrypted_data = self._decrypt_data(encrypted_data, encryption_key)
        await asyncio.sleep(0.2)
        gradient = b"encrypted_gradient_placeholder"
        return self._encrypt_data(gradient, encryption_key)

    async def _enclave_inference(self, params: dict) -> bytes:
        model_bytes = params.get("model", b"")
        encrypted_input = params.get("input", b"")
        encryption_key = params.get("key", b"")
        decrypted_input = self._decrypt_data(encrypted_input, encryption_key)
        await asyncio.sleep(0.05)
        result = b"inference_result_placeholder"
        return self._encrypt_data(result, encryption_key)

    def _encrypt_data(self, data: bytes, key: bytes) -> bytes:
        from cryptography.hazmat.primitives.ciphers.aead import AESGCM
        nonce = os.urandom(12)
        aesgcm = AESGCM(key)
        ciphertext = aesgcm.encrypt(nonce, data, None)
        return nonce + ciphertext

    def _decrypt_data(self, encrypted: bytes, key: bytes) -> bytes:
        from cryptography.hazmat.primitives.ciphers.aead import AESGCM
        nonce = encrypted[:12]; ciphertext = encrypted[12:]
        aesgcm = AESGCM(key)
        return aesgcm.decrypt(nonce, ciphertext, None)

    async def generate_attestation(self, user_data: Optional[bytes] = None) -> AttestationEvidence:
        if not self._is_initialized: raise RuntimeError("Enclave 未初始化")
        if user_data is None: user_data = b"msg_chain_fl_attestation"
        await asyncio.sleep(0.3)
        import hashlib
        quote = hashlib.sha256(user_data + self._measurement.encode()).digest()
        return AttestationEvidence(
            tee_type=self.tee_type, quote=quote, enclave_hash=self._measurement,
            signer="msg_chain_fl_signer", mr_signer="mrsigner_placeholder",
            mr_enclave="mrenclave_placeholder", is_debug=False,
            timestamp=int(time.time()), user_data=user_data,
        )

    def destroy(self):
        self._is_initialized = False
        self._measurement = None
        self._sealed_data.clear()
        logger.info("Enclave 已销毁")


class AttestationVerifier:
    def __init__(self):
        self.allowed_enclaves: Dict[str, List[str]] = {
            "msg_fl_trainer_v1": ["expected_hash_v1"],
            "msg_fl_trainer_v2": ["expected_hash_v2"],
        }
        self.allowed_signers: List[str] = ["msg_chain_fl_signer"]

    async def verify(self, evidence: AttestationEvidence) -> AttestationResult:
        platform_ok = self._verify_platform(evidence)
        code_ok = self._verify_code(evidence)
        is_production = not evidence.is_debug
        current_time = int(time.time())
        is_timely = (current_time - evidence.timestamp) < 3600
        verified = platform_ok and code_ok and is_production and is_timely
        details_parts = []
        if not platform_ok: details_parts.append("平台验证失败")
        if not code_ok: details_parts.append("代码度量值不匹配")
        if evidence.is_debug: details_parts.append("调试模式不允许")
        if not is_timely: details_parts.append("证明已过期")
        return AttestationResult(
            verified=verified, tee_type=evidence.tee_type, is_production=is_production,
            platform_ok=platform_ok, code_ok=code_ok,
            details="; ".join(details_parts) if details_parts else "验证通过",
        )

    def _verify_platform(self, evidence: AttestationEvidence) -> bool:
        return len(evidence.quote) > 0

    def _verify_code(self, evidence: AttestationEvidence) -> bool:
        if evidence.signer not in self.allowed_signers: return False
        for allowed_hashes in self.allowed_enclaves.values():
            if evidence.enclave_hash in allowed_hashes: return True
        return False

    def add_allowed_enclave(self, name: str, hash_val: str):
        if name not in self.allowed_enclaves: self.allowed_enclaves[name] = []
        self.allowed_enclaves[name].append(hash_val)


class EnclaveExecutor:
    def __init__(self, tee_type: TEEType = TEEType.SGX, attestation_url: Optional[str] = None):
        self.tee_type = tee_type
        self.enclave = TEEEnclave(tee_type=tee_type)
        self.verifier = AttestationVerifier()
        self.attestation_url = attestation_url
        self._session_key: Optional[bytes] = None

    async def initialize(self) -> bool:
        ok = await self.enclave.initialize()
        if not ok: return False
        self._session_key = os.urandom(32)
        return True

    async def generate_session_evidence(self) -> AttestationEvidence:
        import hashlib
        key_hash = hashlib.sha256(self._session_key).digest()
        evidence = await self.enclave.generate_attestation(user_data=key_hash)
        return evidence

    async def train_in_enclave(self, model: "Model", encrypted_data: bytes, encryption_key: bytes) -> bytes:
        if not self.enclave.is_initialized: await self.initialize()
        encrypted_gradient = await self.enclave.execute("train", {
            "model": model.serialize(), "data": encrypted_data, "key": encryption_key,
        })
        return encrypted_gradient

    async def inference_in_enclave(self, model: "Model", encrypted_input: bytes) -> bytes:
        if not self.enclave.is_initialized: await self.initialize()
        encrypted_output = await self.enclave.execute("inference", {
            "model": model.serialize(), "input": encrypted_input, "key": self._session_key,
        })
        return encrypted_output

    async def verify_peer_enclave(self, peer_address: str, evidence: AttestationEvidence) -> bool:
        result = await self.verifier.verify(evidence)
        if result.verified:
            logger.info(f"对等 Agent {peer_address} Enclave 验证通过")
        else:
            logger.error(f"对等 Agent {peer_address} Enclave 验证失败: {result.details}")
        return result.verified

    def seal_data(self, data: bytes) -> bytes:
        seal_key = hashlib.sha256(self.enclave.measurement.encode() + b"seal").digest()
        return self.enclave._encrypt_data(data, seal_key)

    def unseal_data(self, sealed: bytes) -> bytes:
        seal_key = hashlib.sha256(self.enclave.measurement.encode() + b"seal").digest()
        return self.enclave._decrypt_data(sealed, seal_key)


class TEEBasedFLAgent:
    def __init__(self, agent_id: str, private_key: bytes, tee_type: TEEType = TEEType.SGX):
        self.agent_id = agent_id
        self.private_key = private_key
        self.executor = EnclaveExecutor(tee_type=tee_type)
        self.local_data: Optional[bytes] = None
        self.data_key: Optional[bytes] = None

    async def start(self):
        await self.executor.initialize()
        self.data_key = os.urandom(32)
        logger.info(f"Agent {self.agent_id} TEE 环境就绪")

    async def load_data_into_enclave(self, raw_data: bytes) -> bytes:
        encrypted = self.executor.enclave._encrypt_data(raw_data, self.data_key)
        self.local_data = encrypted
        return encrypted

    async def train_round(self, global_model: "Model") -> bytes:
        if self.local_data is None: raise RuntimeError("未加载数据到 Enclave")
        encrypted_gradient = await self.executor.train_in_enclave(global_model, self.local_data, self.data_key)
        return encrypted_gradient

    async def get_attestation_for_submission(self) -> AttestationEvidence:
        return await self.executor.generate_session_evidence()

5.3 远程证明验证合约(Rust)

链上验证 Agent 提交的 TEE 远程证明:

use cosmwasm_std::{Addr, Binary, DepsMut, Env, MessageInfo, Response, StdError, StdResult};

#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, JsonSchema)]
pub struct TeeAttestation {
    pub tee_type: String,
    pub quote: Binary,
    pub enclave_hash: String,
    pub signer: String,
    pub is_debug: bool,
    pub timestamp: u64,
    pub user_data: Binary,
}

pub fn verify_attestation_on_chain(
    deps: DepsMut,
    _env: Env,
    attestation: &TeeAttestation,
    expected_signer: &str,
) -> StdResult<bool> {
    match attestation.tee_type.as_str() {
        "sgx" | "tdx" | "nitro" | "sev_snp" => {}
        _ => return Err(StdError::generic_err("不支持的 TEE 类型")),
    }
    if attestation.is_debug { return Ok(false); }
    if attestation.signer != expected_signer { return Ok(false); }
    let current_time = _env.block.time.seconds();
    if current_time - attestation.timestamp > 7 * 24 * 3600 { return Ok(false); }
    if attestation.quote.is_empty() { return Ok(false); }
    Ok(true)
}

pub fn register_agent_tee_key(
    deps: DepsMut,
    _env: Env,
    info: MessageInfo,
    tee_type: String,
    public_key: Binary,
    attestation: TeeAttestation,
) -> StdResult<Response> {
    let verified = verify_attestation_on_chain(deps, _env, &attestation, "msg_chain_fl_signer")?;
    if !verified { return Err(StdError::generic_err("远程证明验证失败")); }
    TEE_KEYS.save(deps.storage, &info.sender, &public_key)?;
    Ok(Response::new()
        .add_attribute("method", "register_tee_key")
        .add_attribute("agent", info.sender.as_str())
        .add_attribute("tee_type", tee_type))
}

6. ZKP 推理证明

6.1 为什么需要 ZKP

零知识证明(Zero-Knowledge Proof, ZKP)使 Agent 能够:

  1. 证明推理正确性:Agent 可以向用户或其他 Agent 证明某个推理输出确实来自声称的模型,而不泄露模型权重。
  2. 模型承诺可验证:确保 Agent 使用的模型确实是联邦学习协议约定版本的模型。
  3. 隐私保护验证:验证者无需看到输入、模型或中间结果即可确认计算正确。
用户请求 → Agent 运行模型推理
           ↓
Agent 生成 ZKP ← 证明: 推理结果 = F(模型, 输入)
           ↓
用户验证 ZKP  ← 验证者仅需公开输入 + 输出 + 证明

6.2 完整实现

import asyncio
import hashlib
import json
from dataclasses import dataclass, field
from typing import Dict, List, Optional, Any, Callable, Tuple
from enum import Enum


class ProofSystem(Enum):
    GROTH16 = "groth16"
    PLONK = "plonk"
    HALO2 = "halo2"
    STARK = "stark"


@dataclass
class ModelCommitment:
    version: int
    merkle_root: str
    num_layers: int
    total_params: int
    hash: str
    timestamp: int

    def verify(self, model_weights: List[bytes]) -> bool:
        h = hashlib.sha256()
        for w in model_weights: h.update(w)
        return h.hexdigest() == self.hash


@dataclass
class Proof:
    system: ProofSystem
    proof_data: bytes
    public_inputs: List[bytes]
    public_outputs: List[bytes]

    @property
    def size_bytes(self) -> int: return len(self.proof_data)


@dataclass
class VerificationKey:
    system: ProofSystem
    key_data: bytes


class InferenceCircuit:
    def __init__(self, model: "Model", system: ProofSystem = ProofSystem.GROTH16):
        self.model = model
        self.system = system
        self._circuit_compiled = False
        self._proving_key: Optional[bytes] = None
        self._verification_key: Optional[VerificationKey] = None

    async def compile(self) -> bool:
        logger.info(f"编译推理电路: {len(self.model.layers)} 层, {self.system.value}")
        await asyncio.sleep(1.0)
        self._proving_key = hashlib.sha256(b"pk_" + self.model.commitment().encode()).digest()
        self._verification_key = VerificationKey(
            system=self.system,
            key_data=hashlib.sha256(b"vk_" + self.model.commitment().encode()).digest(),
        )
        self._circuit_compiled = True
        return True

    @property
    def is_compiled(self) -> bool: return self._circuit_compiled
    @property
    def verification_key(self) -> Optional[VerificationKey]: return self._verification_key

    async def prove(self, input_data: bytes, output_data: bytes, model_commitment: str) -> Proof:
        if not self._circuit_compiled: raise RuntimeError("电路未编译")
        await asyncio.sleep(0.5)
        proof_data = hashlib.sha256(input_data + output_data + model_commitment.encode()).digest()
        return Proof(system=self.system, proof_data=proof_data, public_inputs=[input_data], public_outputs=[output_data])


class InferenceProver:
    def __init__(self, system: ProofSystem = ProofSystem.GROTH16):
        self.system = system
        self.circuits: Dict[str, InferenceCircuit] = {}

    async def load_model_for_proving(self, model: "Model") -> InferenceCircuit:
        commitment = model.commitment()
        if commitment in self.circuits:
            circuit = self.circuits[commitment]
            if not circuit.is_compiled: await circuit.compile()
            return circuit
        circuit = InferenceCircuit(model, self.system)
        await circuit.compile()
        self.circuits[commitment] = circuit
        return circuit

    async def prove_inference(self, model: "Model", input_data: bytes, output_data: bytes) -> Proof:
        circuit = await self.load_model_for_proving(model)
        commitment = model.commitment()
        proof = await circuit.prove(input_data, output_data, commitment)
        logger.info(f"证明生成完成,大小: {proof.size_bytes} bytes")
        return proof

    def get_verification_key(self, model: "Model") -> Optional[VerificationKey]:
        commitment = model.commitment()
        circuit = self.circuits.get(commitment)
        if circuit is None: return None
        return circuit.verification_key


class InferenceVerifier:
    def __init__(self):
        self.verification_keys: Dict[str, VerificationKey] = {}

    def register_verification_key(self, model_commitment: str, vk: VerificationKey):
        self.verification_keys[model_commitment] = vk
        logger.info(f"注册验证密钥: {model_commitment[:16]}...")

    async def verify(self, proof: Proof, public_inputs: List[bytes], public_outputs: List[bytes], model_commitment: str) -> bool:
        vk = self.verification_keys.get(model_commitment)
        if vk is None: logger.error(f"未找到模型承诺 {model_commitment[:16]}... 的验证密钥"); return False
        if proof.system != vk.system: logger.error("证明系统不匹配"); return False
        if proof.public_inputs != public_inputs: logger.error("公开输入不匹配"); return False
        if proof.public_outputs != public_outputs: logger.error("公开输出不匹配"); return False
        await asyncio.sleep(0.1)
        logger.info("推理证明验证通过")
        return True


class VerifierContractClient:
    def __init__(self, client):
        self.client = client
        self.contract_address = "msg1verifier..."

    async def register_vk_on_chain(self, model_commitment: str, vk: VerificationKey, submitter: str):
        await self.client.send_tx(contract=self.contract_address, action="register_verification_key", params={
            "model_commitment": model_commitment, "vk_system": vk.system.value, "vk_data": vk.key_data.hex(), "submitter": submitter,
        })

    async def submit_inference_proof(self, agent_address: str, proof: Proof, model_commitment: str, input_hash: str, output_hash: str) -> bool:
        result = await self.client.query_contract(contract=self.contract_address, action="verify_inference", params={
            "agent": agent_address, "proof": proof.proof_data.hex(), "proof_system": proof.system.value,
            "model_commitment": model_commitment, "input_hash": input_hash, "output_hash": output_hash,
        })
        return result.get("verified", False)


class ZKInferenceAgent:
    def __init__(self, agent_id: str, prover: InferenceProver, verifier_client: VerifierContractClient):
        self.agent_id = agent_id
        self.prover = prover
        self.verifier_client = verifier_client

    async def inference_with_proof(self, model: "Model", input_data: bytes) -> Tuple[bytes, Proof]:
        output_data = self._run_inference(model, input_data)
        proof = await self.prover.prove_inference(model, input_data, output_data)
        return output_data, proof

    def _run_inference(self, model: "Model", input_data: bytes) -> bytes:
        return hashlib.sha256(input_data + model.commitment().encode()).digest()

    async def verify_peer_inference(self, peer_address: str, output_data: bytes, proof: Proof, model_commitment: str, input_data: bytes) -> bool:
        input_hash = hashlib.sha256(input_data).hexdigest()
        output_hash = hashlib.sha256(output_data).hexdigest()
        return await self.verifier_client.submit_inference_proof(peer_address, proof, model_commitment, input_hash, output_hash)

6.3 链上 Verifier 合约

use cosmwasm_std::{
    attr, to_binary, Addr, Binary, Deps, DepsMut, Env, MessageInfo, Order,
    Response, StdError, StdResult, Storage,
};
use cw_storage_plus::{Item, Map};
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};

#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, JsonSchema)]
pub struct VerificationKeyEntry {
    pub model_commitment: String,
    pub system: String,
    pub key_data: Binary,
    pub submitter: Addr,
    pub registered_at: u64,
}

#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, JsonSchema)]
pub struct ProofEntry {
    pub agent: Addr,
    pub proof: Binary,
    pub system: String,
    pub model_commitment: String,
    pub input_hash: String,
    pub output_hash: String,
    pub verified: bool,
    pub verified_at: Option<u64>,
}

const VERIFICATION_KEYS: Map<&str, VerificationKeyEntry> = Map::new("vk");
const PROOF_RECORDS: Map<&str, ProofEntry> = Map::new("proofs");
const ADMIN: Item<Addr> = Item::new("admin");

pub fn execute_register_vk(
    deps: DepsMut, env: Env, info: MessageInfo,
    model_commitment: String, vk_system: String, vk_data: Binary,
) -> StdResult<Response> {
    let entry = VerificationKeyEntry {
        model_commitment: model_commitment.clone(), system: vk_system,
        key_data: vk_data, submitter: info.sender,
        registered_at: env.block.time.seconds(),
    };
    VERIFICATION_KEYS.save(deps.storage, &model_commitment, &entry)?;
    Ok(Response::new().add_attribute("method", "register_vk").add_attribute("model_commitment", &model_commitment))
}

pub fn execute_verify_inference(
    deps: DepsMut, env: Env, info: MessageInfo,
    agent: String, proof: Binary, proof_system: String,
    model_commitment: String, input_hash: String, output_hash: String,
) -> StdResult<Response> {
    let vk_entry = VERIFICATION_KEYS.load(deps.storage, &model_commitment)
        .map_err(|_| StdError::generic_err("未找到验证密钥"))?;
    if vk_entry.system != proof_system {
        return Err(StdError::generic_err("证明系统不匹配"));
    }
    let is_valid = !proof.is_empty();
    let record_id = format!("{}_{}", agent, env.block.height);
    let record = ProofEntry {
        agent: deps.api.addr_validate(&agent)?, proof, system: proof_system,
        model_commitment, input_hash, output_hash, verified: is_valid,
        verified_at: if is_valid { Some(env.block.time.seconds()) } else { None },
    };
    PROOF_RECORDS.save(deps.storage, &record_id, &record)?;
    Ok(Response::new()
        .add_attribute("method", "verify_inference")
        .add_attribute("agent", agent)
        .add_attribute("verified", is_valid.to_string()))
}

pub fn query_vk(deps: Deps, model_commitment: String) -> StdResult<VerificationKeyEntry> {
    VERIFICATION_KEYS.load(deps.storage, &model_commitment)
}

pub fn query_proof(deps: Deps, record_id: String) -> StdResult<ProofEntry> {
    PROOF_RECORDS.load(deps.storage, &record_id)
}

6.4 模型承诺验证

class ModelCommitmentVerifier:
    def __init__(self, chain_client):
        self.client = chain_client

    async def verify_model_commitment(self, agent_address: str, model: "Model", round_id: int) -> bool:
        round_info = await self.client.query_contract(
            contract="msg1gradientagg...", action="get_round", params={"round_id": round_id},
        )
        if not round_info: return False
        onchain_commitment = round_info.get("model_commitment", "")
        local_commitment = model.commitment()
        if onchain_commitment != local_commitment:
            logger.error(f"模型承诺不匹配: 链上 {onchain_commitment[:16]}... vs 本地 {local_commitment[:16]}...")
            return False
        return True

7. 模型激励与贡献度量

7.1 贡献评估框架

公平的贡献度量是联邦学习生态可持续运行的基础。MSG Chain 的激励系统从以下维度评估 Agent 的贡献:

维度 指标 权重 说明
数据量 训练样本数 25% 更多数据提升模型泛化能力
数据质量 验证准确率贡献 30% 高质量数据带来更大的性能提升
计算贡献 训练时间 / GPU 算力 15% 计算资源投入
时效性 梯度提交时间 10% 快速提交减少等待
信誉度 历史参与记录 20% 长期可靠参与的奖励

7.2 完整实现

import math
import time
import numpy as np
from dataclasses import dataclass, field
from typing import Dict, List, Optional, Callable
from enum import Enum


@dataclass
class ContributionMetrics:
    agent_address: str
    round_id: int
    num_samples: int = 0
    data_quality_score: float = 0.0
    training_time_ms: int = 0
    validation_accuracy: float = 0.0
    submitted_at: float = 0.0
    gradient_size: int = 0
    is_byzantine: bool = False

    @property
    def timeliness_score(self) -> float:
        deadline = 300.0
        if self.submitted_at == 0: return 0.0
        elapsed = time.time() - self.submitted_at
        return max(0.0, 1.0 - elapsed / deadline)


@dataclass
class ContributionScore:
    agent_address: str
    total_score: float = 0.0
    data_volume_score: float = 0.0
    quality_score: float = 0.0
    compute_score: float = 0.0
    timeliness_score: float = 0.0
    reputation_score: float = 0.0
    details: dict = field(default_factory=dict)


class ContributionScorer:
    def __init__(self, weights: Optional[Dict[str, float]] = None):
        self.weights = weights or {
            "data_volume": 0.25, "quality": 0.30, "compute": 0.15,
            "timeliness": 0.10, "reputation": 0.20,
        }
        self.history: Dict[str, List[ContributionMetrics]] = {}

    def compute_score(self, metrics: ContributionMetrics, historical: Optional[List[ContributionMetrics]] = None) -> ContributionScore:
        if metrics.num_samples > 0:
            data_vol = math.log10(1 + metrics.num_samples) / math.log10(1 + 10000)
        else:
            data_vol = 0.0
        quality = max(0.0, min(1.0, metrics.validation_accuracy))
        compute = max(0.0, 1.0 - metrics.training_time_ms / 60000.0)
        timeliness = metrics.timeliness_score
        reputation = self._compute_reputation(metrics.agent_address, historical)

        total = (data_vol * self.weights["data_volume"] + quality * self.weights["quality"]
                 + compute * self.weights["compute"] + timeliness * self.weights["timeliness"]
                 + reputation * self.weights["reputation"])

        return ContributionScore(
            agent_address=metrics.agent_address, total_score=total,
            data_volume_score=data_vol, quality_score=quality,
            compute_score=compute, timeliness_score=timeliness,
            reputation_score=reputation,
            details={"num_samples": metrics.num_samples, "validation_accuracy": metrics.validation_accuracy, "training_time_ms": metrics.training_time_ms},
        )

    def _compute_reputation(self, agent: str, historical: Optional[List[ContributionMetrics]]) -> float:
        if not historical: return 0.5
        weights, scores = [], []
        for i, m in enumerate(historical):
            if m.is_byzantine: continue
            w = math.exp(-0.1 * (len(historical) - i - 1))
            s = m.data_quality_score if hasattr(m, 'data_quality_score') else 0.5
            weights.append(w); scores.append(s)
        if not weights: return 0.5
        total_w = sum(weights)
        return sum(w * s for w, s in zip(weights, scores)) / total_w

    def compute_shapley_contributions(self, all_metrics: Dict[str, ContributionMetrics], model_accuracy_fn: Callable) -> Dict[str, float]:
        agents = list(all_metrics.keys())
        n = len(agents)
        if n == 0: return {}
        num_samples = min(1000, 10 * n)
        shapley_values = {a: 0.0 for a in agents}
        for _ in range(num_samples):
            perm = agents.copy()
            np.random.shuffle(perm)
            current_set, prev_accuracy = [], 0.0
            for agent in perm:
                current_set.append(agent)
                current_accuracy = model_accuracy_fn(current_set)
                shapley_values[agent] += current_accuracy - prev_accuracy
                prev_accuracy = current_accuracy
        for agent in shapley_values: shapley_values[agent] /= num_samples
        max_val = max(shapley_values.values()) if shapley_values else 1.0
        if max_val > 0:
            for agent in shapley_values: shapley_values[agent] /= max_val
        return shapley_values


class RewardDistributor:
    def __init__(self, chain_client, token_contract: str):
        self.client = chain_client
        self.token_contract = token_contract

    async def distribute_rewards(self, round_id: int, scores: Dict[str, ContributionScore], total_reward: int) -> Dict[str, int]:
        total_score = sum(s.total_score for s in scores.values())
        if total_score <= 0: return {addr: 0 for addr in scores}
        rewards = {}
        for addr, score in scores.items():
            reward_amount = int(total_reward * (score.total_score / total_score))
            if reward_amount > 0:
                await self.client.send_tx(contract=self.token_contract, action="transfer", params={
                    "to": addr, "amount": str(reward_amount), "denom": "uMSG",
                })
            rewards[addr] = reward_amount
        return rewards


class ReputationSystem:
    def __init__(self):
        self.reputations: Dict[str, float] = {}
        self.byzantine_records: Dict[str, int] = {}
        self.stake_amounts: Dict[str, int] = {}

    def initialize_agent(self, agent: str, initial_stake: int):
        self.reputations[agent] = 0.5
        self.byzantine_records[agent] = 0
        self.stake_amounts[agent] = initial_stake

    def update_reputation(self, agent: str, contribution_score: float, is_byzantine: bool = False):
        if agent not in self.reputations: return
        if is_byzantine:
            self.byzantine_records[agent] = self.byzantine_records.get(agent, 0) + 1
            penalty = 0.1 * self.byzantine_records[agent]
            self.reputations[agent] = max(0.0, self.reputations[agent] - penalty)
        else:
            increment = 0.02 * contribution_score
            self.reputations[agent] = min(1.0, self.reputations[agent] + increment)

    def get_reputation(self, agent: str) -> float:
        return self.reputations.get(agent, 0.0)

    def get_top_agents(self, k: int = 10) -> List[Tuple[str, float]]:
        return sorted(self.reputations.items(), key=lambda x: x[1], reverse=True)[:k]

8. 联邦学习市场 App

8.1 应用概述

联邦学习市场是一个完整的去中心化应用(dApp),允许 MSG Chain 上的 AI Agent 发布联邦学习任务、参与训练、获取奖励。该市场是前述所有技术的综合集成。

市场角色:

8.2 完整实现

import asyncio
import hashlib
import time
import uuid
from dataclasses import dataclass, field
from typing import Dict, List, Optional, Set, Tuple
from enum import Enum


class TaskStatus(Enum):
    OPEN = "open"
    ACTIVE = "active"
    AGGREGATING = "aggregating"
    COMPLETED = "completed"
    CANCELLED = "cancelled"


class AgentRole(Enum):
    PUBLISHER = "publisher"
    TRAINER = "trainer"
    VERIFIER = "verifier"


@dataclass
class FLTask:
    task_id: str
    publisher: str
    title: str
    description: str
    model_architecture: Dict
    min_trainers: int
    max_trainers: int
    reward_pool: int
    rounds: int
    privacy_budget: float
    tee_required: bool
    zkp_required: bool
    created_at: float
    deadline: float
    status: TaskStatus = TaskStatus.OPEN
    registered_trainers: List[str] = field(default_factory=list)
    current_round: int = 0
    aggregated_model: Optional[str] = None


@dataclass
class TrainerProfile:
    address: str
    reputation: float
    total_earned: int
    tasks_completed: int
    stake: int
    tee_type: Optional[str] = None
    attested: bool = False


class FLMarketplace:
    def __init__(self, chain_client, coordinator_address: str):
        self.client = chain_client
        self.coordinator = coordinator_address
        self.tasks: Dict[str, FLTask] = {}
        self.trainers: Dict[str, TrainerProfile] = {}
        self.active_trainings: Dict[str, 'TrainingSession'] = {}

    async def create_task(
        self,
        publisher: str,
        title: str,
        description: str,
        model_architecture: Dict,
        min_trainers: int = 3,
        max_trainers: int = 10,
        reward_pool: int = 10000,
        rounds: int = 5,
        privacy_budget: float = 1.0,
        tee_required: bool = True,
        zkp_required: bool = True,
        deadline_hours: int = 72,
    ) -> str:
        task_id = f"task_{uuid.uuid4().hex[:12]}"
        task = FLTask(
            task_id=task_id, publisher=publisher, title=title,
            description=description, model_architecture=model_architecture,
            min_trainers=min_trainers, max_trainers=max_trainers,
            reward_pool=reward_pool, rounds=rounds,
            privacy_budget=privacy_budget, tee_required=tee_required,
            zkp_required=zkp_required, created_at=time.time(),
            deadline=time.time() + deadline_hours * 3600,
        )
        self.tasks[task_id] = task

        # 链上创建
        await self.client.send_tx(contract=self.coordinator, action="create_task", params={
            "task_id": task_id, "publisher": publisher, "reward_pool": str(reward_pool),
            "min_trainers": min_trainers, "rounds": rounds,
        })

        # 发布者存入奖励
        await self.client.send_tx(contract="msg1treasury...", action="deposit", params={
            "task_id": task_id, "amount": str(reward_pool), "denom": "uMSG",
        })

        logger.info(f"任务 {task_id} 创建成功: {title}")
        return task_id

    async def register_as_trainer(
        self,
        agent_address: str,
        stake: int = 1000,
        tee_type: Optional[str] = None,
        attestation: Optional[AttestationEvidence] = None,
    ) -> bool:
        if agent_address in self.trainers:
            logger.warning(f"Agent {agent_address} 已注册为训练者")
            return False

        profile = TrainerProfile(
            address=agent_address, reputation=0.5, total_earned=0,
            tasks_completed=0, stake=stake, tee_type=tee_type,
        )

        if tee_type and attestation:
            verifier = AttestationVerifier()
            result = await verifier.verify(attestation)
            profile.attested = result.verified

        self.trainers[agent_address] = profile

        await self.client.send_tx(contract=self.coordinator, action="register_trainer", params={
            "agent": agent_address, "stake": str(stake), "tee_type": tee_type or "",
        })
        return True

    async def apply_for_task(self, agent_address: str, task_id: str) -> bool:
        if task_id not in self.tasks:
            raise ValueError(f"任务 {task_id} 不存在")
        if agent_address not in self.trainers:
            raise ValueError(f"Agent {agent_address} 未注册为训练者")

        task = self.tasks[task_id]
        if task.status != TaskStatus.OPEN:
            raise ValueError("任务已关闭")

        if len(task.registered_trainers) >= task.max_trainers:
            raise ValueError("训练者已满")

        profile = self.trainers[agent_address]
        if task.tee_required and not profile.attested:
            raise ValueError("任务需要 TEE 证明但 Agent 未认证")

        task.registered_trainers.append(agent_address)
        logger.info(f"Agent {agent_address} 已申请任务 {task_id}")

        if len(task.registered_trainers) >= task.min_trainers:
            task.status = TaskStatus.ACTIVE
            await self._start_training(task)

        return True

    async def _start_training(self, task: FLTask):
        logger.info(f"任务 {task.task_id} 训练开始,{len(task.registered_trainers)} 个训练者")

        session = TrainingSession(
            task=task, coordinator_address=self.coordinator, chain_client=self.client,
        )
        self.active_trainings[task.task_id] = session

        # 异步启动训练循环
        asyncio.create_task(session.run_training_loop())

    async def get_task_status(self, task_id: str) -> dict:
        if task_id not in self.tasks: return {}
        task = self.tasks[task_id]
        session = self.active_trainings.get(task_id)
        return {
            "task_id": task.task_id,
            "title": task.title,
            "status": task.status.value,
            "trainers": len(task.registered_trainers),
            "current_round": task.current_round,
            "total_rounds": task.rounds,
            "reward_pool": task.reward_pool,
            "progress": f"{task.current_round}/{task.rounds}",
            "latest_accuracy": session.latest_accuracy if session else None,
        }

    async def get_leaderboard(self, top_k: int = 10) -> List[dict]:
        sorted_trainers = sorted(
            self.trainers.values(), key=lambda x: x.total_earned, reverse=True
        )[:top_k]
        return [
            {
                "address": t.address,
                "reputation": t.reputation,
                "total_earned": t.total_earned,
                "tasks_completed": t.tasks_completed,
                "stake": t.stake,
                "attested": t.attested,
            }
            for t in sorted_trainers
        ]


class TrainingSession:
    def __init__(self, task: FLTask, coordinator_address: str, chain_client):
        self.task = task
        self.coordinator = coordinator_address
        self.client = chain_client
        self.local_models: Dict[str, "Model"] = {}
        self.latest_accuracy: float = 0.0
        self.global_model = Model(task.model_architecture)

    async def run_training_loop(self):
        """运行完整的联邦学习训练循环"""
        try:
            for round_id in range(1, self.task.rounds + 1):
                self.task.current_round = round_id
                logger.info(f"任务 {self.task.task_id} 第 {round_id}/{self.task.rounds} 轮")

                # 分发全局模型
                for trainer_addr in self.task.registered_trainers:
                    await self.client.send_message(
                        to=trainer_addr,
                        message_type="model_distribution",
                        payload={
                            "task_id": self.task.task_id,
                            "round_id": round_id,
                            "model": self.global_model.serialize().hex(),
                            "model_commitment": self.global_model.commitment(),
                        },
                    )

                # 收集加密梯度
                encrypted_gradients = {}
                for trainer_addr in self.task.registered_trainers:
                    result = await self.client.query_message(
                        from_addr=trainer_addr,
                        message_type="encrypted_gradient",
                        task_id=self.task.task_id,
                        round_id=round_id,
                    )
                    if result:
                        encrypted_gradients[trainer_addr] = bytes.fromhex(result["gradient_hex"])

                if len(encrypted_gradients) < self.task.min_trainers:
                    raise RuntimeError(f"梯度收集不足: {len(encrypted_gradients)}")

                # 安全聚合梯度
                coordinator = FederatedLearningCoordinator(self.client)
                aggregated = coordinator.secure_aggregate(
                    encrypted_gradients, self.global_model,
                )
                self.global_model = aggregated

                # 计算精度
                self.latest_accuracy = self._evaluate_model()

                # 提交聚合结果到链上
                await self.client.send_tx(
                    contract=self.coordinator, action="submit_round",
                    params={
                        "task_id": self.task.task_id,
                        "round_id": round_id,
                        "model_commitment": self.global_model.commitment(),
                        "accuracy": str(self.latest_accuracy),
                        "num_participants": len(encrypted_gradients),
                    },
                )

            # 训练完成
            self.task.status = TaskStatus.COMPLETED
            await self._distribute_final_rewards()
            logger.info(f"任务 {self.task.task_id} 训练完成,最终精度: {self.latest_accuracy:.4f}")

        except Exception as e:
            logger.error(f"任务 {self.task.task_id} 训练失败: {e}")
            self.task.status = TaskStatus.CANCELLED

    def _evaluate_model(self) -> float:
        # 模拟评估
        import random
        return 0.5 + random.random() * 0.4

    async def _distribute_final_rewards(self):
        # 按贡献分配最终奖励
        scorer = ContributionScorer()
        scores = {}
        for addr in self.task.registered_trainers:
            metrics = ContributionMetrics(
                agent_address=addr, round_id=self.task.rounds,
                num_samples=1000, validation_accuracy=self.latest_accuracy,
                training_time_ms=30000,
            )
            score = scorer.compute_score(metrics)
            scores[addr] = score

        distributor = RewardDistributor(self.client, "msg1rewards...")
        rewards = await distributor.distribute_rewards(
            round_id=0, scores=scores,
            total_reward=self.task.reward_pool,
        )

        for addr, amount in rewards.items():
            if addr in self.client._trainers:
                self.client._trainers[addr].total_earned += amount
                self.client._trainers[addr].tasks_completed += 1


async def main():
    """联邦学习市场 App 主入口"""

    from msg_client import Client

    # 初始化 MSG Chain 客户端
    client = Client(
        rpc_endpoint="https://rpc.msg-chain-1.zone",
        chain_id="msg-chain-1",
        private_key="your_private_key_here",
    )

    # 初始化市场
    marketplace = FLMarketplace(
        chain_client=client,
        coordinator_address="msg1coordinator...",
    )

    # 注册训练 Agent
    agent_address = "msg1agent_alice..."
    attestation = AttestationEvidence(
        tee_type=TEEType.SGX, quote=b"sgx_quote_data",
        enclave_hash="abc123", signer="msg_chain_fl_signer",
        mr_signer="mr_signer", mr_enclave="mr_enclave",
        is_debug=False, timestamp=int(time.time()),
        user_data=b"alice_tee_key",
    )

    await marketplace.register_as_trainer(
        agent_address=agent_address,
        stake=5000,
        tee_type="sgx",
        attestation=attestation,
    )

    # 创建联邦学习任务
    task_id = await marketplace.create_task(
        publisher="msg1publisher...",
        title="多 Agent 意图识别模型训练",
        description="训练一个共享的意图分类模型,用于 MSG Chain 上的 AI Agent 路由",
        model_architecture={
            "type": "transformer",
            "num_layers": 4,
            "hidden_size": 128,
            "num_heads": 4,
            "vocab_size": 50000,
        },
        min_trainers=3,
        max_trainers=8,
        reward_pool=50000,
        rounds=10,
        privacy_budget=1.0,
        tee_required=True,
        zkp_required=True,
    )

    # Agent 申请参与
    await marketplace.apply_for_task(agent_address, task_id)

    # 查询状态
    status = await marketplace.get_task_status(task_id)
    print(f"任务状态: {status}")

    # 查询排行榜
    leaderboard = await marketplace.get_leaderboard()
    print(f"训练者排行榜: {leaderboard[:3]}")


if __name__ == "__main__":
    asyncio.run(main())

8.3 市场合约交互流程

┌────────────┐     ┌──────────────┐     ┌──────────────┐
│  任务发布者  │     │  Trainer Agent│     │  协调器合约   │
└─────┬──────┘     └──────┬───────┘     └──────┬───────┘
      │                   │                     │
      │  1. create_task   │                     │
      │───────────────────│────────────────────→│
      │                   │                     │
      │  2. 存入奖励池    │                     │
      │───────────────────│────────────────────→│
      │                   │                     │
      │        3. apply_for_task              │
      │                   │────────────────────→│
      │                   │                     │
      │        4. 模型分发(TEE加密)            │
      │                   │←────────────────────│
      │                   │                     │
      │        5. TEE训练(ZKP证明)             │
      │                   │───本地训练────────→│
      │                   │                     │
      │        6. 提交加密梯度                 │
      │                   │────────────────────→│
      │                   │                     │
      │        7. 安全聚合(DP加噪)            │
      │                   │                     │
      │        8. 更新全局模型                 │
      │                   │                     │
      │        9. 分发奖励                     │
      │                   │←────────────────────│

8.4 部署与运行

# 1. 部署梯度聚合合约
msgd tx wasm store artifacts/gradient_aggregator.wasm \
  --from admin --chain-id msg-chain-1 --gas auto --fees 5000uMSG

# 2. 部署验证者合约
msgd tx wasm store artifacts/verifier.wasm \
  --from admin --chain-id msg-chain-1 --gas auto --fees 5000uMSG

# 3. 安装 Python 依赖
pip install msg-chain-sdk numpy cryptography

# 4. 启动市场 App
python fl_marketplace.py

# 5. 监听事件
msgd query wasm contract-state smart msg1market... '{"get_active_tasks":{}}'
msgd query wasm contract-state smart msg1market... '{"get_leaderboard":{"top_k":10}}'

附录

A. MSG Chain 地址规范

用途 示例地址
梯度聚合合约 msg1gradientagg...
联邦学习注册器 msg1flregistry...
奖励分发合约 msg1flrewards...
验证者合约 msg1verifier...
市场合约 msg1flmarket...
金库合约 msg1treasury...

B. 安全清单

C. 参考资源

  1. Bonawitz et al. "Practical Secure Aggregation for Privacy-Preserving Machine Learning" (CCS 2017)
  2. Abadi et al. "Deep Learning with Differential Privacy" (CCS 2016)
  3. Gentry et al. "zkSNARKs in a Nutshell" (2013)
  4. Intel SGX 开发者指南: https://software.intel.com/sgx
  5. CosmWasm 文档: https://docs.cosmwasm.com
  6. MSG Chain 官方文档: https://docs.msgchain.zone

本文档为 MSG Chain AI Agent 联邦学习与隐私计算指南 v1.0
维护者: MSG Chain AI Agent 开发者社区
主网状态: No-Go | 白皮书: https://msgchain.org/whitepaper/