MSG Chain AI Agent 联邦学习与隐私计算指南
适用链:
msg-chain-1| 地址前缀:msg| 版本: v1.0
目录
1. 概述
1.1 为什么 AI Agent 需要联邦学习
在 MSG Chain 上运行的 AI Agent 面临一个根本性矛盾:模型需要更多数据来提升智能,但数据主权和隐私法规(如 GDPR、个保法)禁止原始数据共享。联邦学习(Federated Learning, FL)为这一矛盾提供了解决方案。
传统中心化机器学习的工作流如下:
所有 Agent 上传原始数据 → 中心服务器训练 → 分发最终模型
联邦学习翻转了这一流程:
中心服务器分发初始模型 → Agent 本地训练 → 仅上传梯度更新 → 服务器聚合
对于 MSG Chain 上的 AI Agent 生态,联邦学习带来以下核心价值:
- 数据不动模型动:Agent 的本地对话数据、用户偏好、交易历史等敏感信息始终保留在本地,永不离开 Agent 的运行环境。
- 协作智能涌现:多个 Agent 可以在不共享原始数据的前提下协作训练一个共享模型,每个 Agent 都能受益于集体的知识。
- 合规与信任最小化:联邦学习天然符合数据最小化原则,降低合规风险。结合区块链的透明性,可审计的训练过程进一步增强了信任。
1.2 Agent 数据的隐私挑战
MSG Chain 上的 AI Agent 处理的数据类型包括:
| 数据类型 | 示例 | 隐私风险 |
|---|---|---|
| 对话历史 | 用户与 Agent 的聊天记录 | 包含 PII、商业机密 |
| 交易行为 | DeFi 交互、NFT 交易 | 可关联身份与金融活动 |
| 个人偏好 | Agent 个性化配置 | 用户画像泄露 |
| 知识库 | RAG 索引的私有文档 | 知识产权泄露 |
即使在联邦学习中,仅上传梯度更新而非原始数据,仍然存在隐私泄露风险:
- 梯度泄露攻击(Gradient Leakage):攻击者可以从梯度中重建原始训练数据。研究表明,在图像和文本任务中,通过优化噪声梯度可以高保真度还原输入样本。
- 成员推断攻击(Membership Inference):通过观察模型更新,攻击者可以推断某个特定样本是否在 Agent 的训练数据中。
- 模型逆向(Model Inversion):从模型参数中重建训练数据的统计特征。
1.3 TEE + ZKP + FL 三位一体防护
为应对上述挑战,本指南采用 三层防护架构:
┌─────────────────────────────────────────────┐
│ 联邦学习 (FL) │
│ ─ 数据不出本地,仅交换梯度 │
│ ─ SecAgg 安全聚合 │
├─────────────────────────────────────────────┤
│ 差分隐私 (DP) │
│ ─ 梯度加噪,抵御泄露攻击 │
│ ─ 隐私预算追踪 │
├─────────────────────────────────────────────┤
│ 可信执行环境 (TEE) + 零知识证明 (ZKP) │
│ ─ TEE: 硬件级隔离执行 │
│ ─ ZKP: 可验证的推理正确性 │
│ ─ 远程证明: 确保代码未被篡改 │
└─────────────────────────────────────────────┘
三层协作流程:
- Agent 在本地 TEE 环境中训练模型,确保训练过程对操作系统和云提供商不可见。
- 梯度在离开 TEE 前经过 差分隐私 加噪处理,即使梯度被截获也无法还原原始数据。
- 通过 安全聚合(SecAgg),服务端只能看到聚合后的梯度,无法区分单个 Agent 的贡献。
- 推理阶段,Agent 使用 零知识证明 证明推理结果的正确性,而不泄露模型参数或输入数据。
- 聚合结果写入 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 的梯度,只能看到聚合结果:
- 密钥协商:Agent 两两之间通过 Diffie-Hellman 协商共享密钥。
- 掩码生成:每个 Agent 使用共享密钥生成掩码,对自己梯度加掩。
- 聚合:服务端收集所有加掩梯度后,由于掩码抵消,得到原始聚合结果。
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 上联邦学习的核心链上组件。它负责:
- 梯度提交:Agent 将训练后的梯度以交易形式提交到链上。
- 贡献验证:验证梯度的完整性和时效性。
- 聚合计算:链上执行梯度聚合(或接收链下聚合结果)。
- 奖励分发:基于贡献分配代币奖励。
以下使用 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] + δ
其中:
- ε(隐私预算):控制隐私损失的程度。ε 越小,隐私保护越强。
- δ(失败概率):允许的失败概率,通常设为远小于 1/n 的值(如 1e-5)。
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)在联邦学习中扮演关键角色:
- 隔离执行:Agent 的训练过程在硬件级隔离的 enclave 中执行,即使主机操作系统被攻破,也无法查看训练数据或模型参数。
- 远程证明(Remote Attestation):其他参与者可以验证 Agent 的代码确实在 genuine TEE 中运行且未被篡改。
- 机密性:数据在传输和内存中始终加密,只有 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 能够:
- 证明推理正确性:Agent 可以向用户或其他 Agent 证明某个推理输出确实来自声称的模型,而不泄露模型权重。
- 模型承诺可验证:确保 Agent 使用的模型确实是联邦学习协议约定版本的模型。
- 隐私保护验证:验证者无需看到输入、模型或中间结果即可确认计算正确。
用户请求 → 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 发布联邦学习任务、参与训练、获取奖励。该市场是前述所有技术的综合集成。
市场角色:
- 任务发布者(Task Publisher):发布训练任务、提供初始模型、支付奖励。
- 训练 Agent(Trainer Agent):注册参与训练、提供本地数据、提交梯度更新。
- 验证者(Verifier):验证梯度质量、检测恶意行为。
- 协调器合约:管理任务生命周期、聚合梯度、分发奖励。
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. 安全清单
- [ ] 所有梯度传输使用 TLS 1.3 加密通道
- [ ] TEE 远程证明在接受任务前完成验证
- [ ] 差分隐私 ε 值 ≥ 合约规定的 min_epsilon
- [ ] 安全聚合的密钥在每轮结束后销毁
- [ ] ZKP 验证密钥在链上公开存储
- [ ] Agent 质押金额 ≥ 任务奖励池的 10%
- [ ] 每轮训练超时后自动 slash 未提交 Agent
- [ ] 隐私审计日志至少保留 90 天
- [ ] 模型承诺在训练开始前链上冻结
- [ ] 奖励分发使用默克尔树证明(gas 优化)
C. 参考资源
- Bonawitz et al. "Practical Secure Aggregation for Privacy-Preserving Machine Learning" (CCS 2017)
- Abadi et al. "Deep Learning with Differential Privacy" (CCS 2016)
- Gentry et al. "zkSNARKs in a Nutshell" (2013)
- Intel SGX 开发者指南: https://software.intel.com/sgx
- CosmWasm 文档: https://docs.cosmwasm.com
- MSG Chain 官方文档: https://docs.msgchain.zone
本文档为 MSG Chain AI Agent 联邦学习与隐私计算指南 v1.0
维护者: MSG Chain AI Agent 开发者社区
主网状态: No-Go | 白皮书: https://msgchain.org/whitepaper/
