โ† Back to Contents

Chapter 06: Security and Attack Mitigation

Defend federated learning systems against adversarial threats and ensure robustness

Threat Model in Federated Learning

Federated learning systems face unique security challenges. Unlike centralized learning where the server controls all data, federated systems must trust clients to honestly participate. This opens several attack vectors that can compromise model integrity, privacy, or availability.

Critical Understanding

Adversaries in federated learning can be clients (sending malicious updates), servers (violating privacy), or external parties (eavesdropping on communication). Defense requires a multi-layered approach.

Data Poisoning Attacks

Label Flipping Attack

Attacker corrupts labels in training data to degrade model performance:

class LabelFlippingAttack:
    """
    Simple data poisoning: flip labels to degrade model

    Example: In MNIST, change all 7s to 1s
    """

    def __init__(self, source_label, target_label):
        self.source_label = source_label
        self.target_label = target_label

    def poison_dataset(self, data, labels):
        """Flip labels for poisoning"""
        poisoned_labels = labels.copy()

        # Flip source_label to target_label
        poisoned_labels[labels == self.source_label] = self.target_label

        return data, poisoned_labels


# Example usage
attack = LabelFlippingAttack(source_label=7, target_label=1)
poisoned_data, poisoned_labels = attack.poison_dataset(train_data, train_labels)

# Model trained on poisoned data will misclassify 7 as 1

Backdoor Attack

Inject a backdoor trigger that causes specific misclassification:

class BackdoorAttack:
    """
    Backdoor attack: Insert trigger pattern to cause targeted misclassification

    Example: Images with a small square in corner are classified as target_class
    """

    def __init__(self, trigger_pattern, target_class):
        self.trigger_pattern = trigger_pattern
        self.target_class = target_class

    def create_backdoor_sample(self, image, original_label):
        """Insert trigger and change label"""

        backdoored_image = image.copy()

        # Add trigger (e.g., 3x3 white square in bottom-right)
        backdoored_image[-3:, -3:] = 1.0  # White trigger

        return backdoored_image, self.target_class

    def poison_dataset(self, data, labels, poison_fraction=0.1):
        """
        Poison a fraction of the dataset with backdoors

        Clean samples remain unchanged
        """
        num_samples = len(data)
        num_poison = int(num_samples * poison_fraction)

        poisoned_data = data.copy()
        poisoned_labels = labels.copy()

        # Randomly select samples to poison
        poison_indices = np.random.choice(num_samples, num_poison, replace=False)

        for idx in poison_indices:
            poisoned_data[idx], poisoned_labels[idx] = self.create_backdoor_sample(
                data[idx], labels[idx]
            )

        return poisoned_data, poisoned_labels


# Backdoored model works normally on clean data
# but misclassifies any input with the trigger pattern

Defense: Anomaly Detection

def detect_poisoned_updates(client_updates, global_model, threshold=3.0):
    """
    Detect poisoned updates using statistical outlier detection

    Args:
        client_updates: List of client model updates
        global_model: Current global model
        threshold: Z-score threshold for outlier detection

    Returns:
        List of trusted update indices
    """

    # Calculate distances from global model
    distances = []
    for update in client_updates:
        dist = np.linalg.norm(update - global_model)
        distances.append(dist)

    # Z-score normalization
    mean_dist = np.mean(distances)
    std_dist = np.std(distances)

    z_scores = [(d - mean_dist) / (std_dist + 1e-10) for d in distances]

    # Filter outliers
    trusted_indices = [i for i, z in enumerate(z_scores) if abs(z) < threshold]

    return trusted_indices

Model Poisoning Attacks

Byzantine Attack

Malicious clients send arbitrary model updates to sabotage aggregation:

class ByzantineAttack:
    """
    Byzantine attack: Send arbitrary malicious updates

    Different strategies for maximum damage
    """

    def generate_random_noise_attack(self, global_model, scale=10.0):
        """Send random noise scaled to disrupt aggregation"""
        return np.random.randn(*global_model.shape) * scale

    def generate_sign_flip_attack(self, honest_update):
        """Flip the sign of honest gradient (opposite direction)"""
        return -honest_update

    def generate_amplification_attack(self, honest_update, factor=10.0):
        """Amplify honest update to dominate aggregation"""
        return honest_update * factor

    def generate_targeted_attack(self, global_model, target_weights):
        """Push model toward specific target weights"""
        return target_weights - global_model

Defense: Byzantine-Robust Aggregation

Use robust aggregation methods from Chapter 4:

class ByzantineRobustServer:
    """
    Server with Byzantine-robust aggregation
    """

    def __init__(self, aggregation_method='krum', byzantine_ratio=0.2):
        self.aggregation_method = aggregation_method
        self.byzantine_ratio = byzantine_ratio

    def robust_aggregate(self, client_updates):
        """
        Aggregate with Byzantine resilience

        Options: median, trimmed_mean, krum, multi_krum
        """

        if self.aggregation_method == 'median':
            return coordinate_wise_median(client_updates)

        elif self.aggregation_method == 'trimmed_mean':
            return trimmed_mean(client_updates, trim_ratio=self.byzantine_ratio)

        elif self.aggregation_method == 'krum':
            num_byzantines = int(len(client_updates) * self.byzantine_ratio)
            return krum(client_updates, num_byzantines)

        elif self.aggregation_method == 'multi_krum':
            num_byzantines = int(len(client_updates) * self.byzantine_ratio)
            m = max(1, len(client_updates) - num_byzantines)
            return multi_krum(client_updates, num_byzantines, m)

        else:
            # Fallback to standard averaging
            return np.mean(client_updates, axis=0)

Privacy Attacks

Gradient Inversion Attack

Reconstruct training data from gradients:

class GradientInversionAttack:
    """
    Attempt to reconstruct training data from gradients

    Based on "Deep Leakage from Gradients" (Zhu et al., 2019)
    """

    def __init__(self, model):
        self.model = model

    def reconstruct_data(self, observed_gradients, num_iterations=1000):
        """
        Reconstruct input data that produces observed gradients

        Args:
            observed_gradients: Gradients received from client
            num_iterations: Optimization iterations

        Returns:
            Reconstructed data (approximate)
        """

        # Initialize random dummy data and labels
        dummy_data = np.random.randn(1, *input_shape)
        dummy_labels = np.random.randint(0, num_classes, size=(1,))

        optimizer = Adam(learning_rate=0.1)

        for iteration in range(num_iterations):
            # Compute gradients of dummy data
            dummy_gradients = self.model.compute_gradients(dummy_data, dummy_labels)

            # Loss: L2 distance between dummy and observed gradients
            gradient_diff = sum([np.sum((dg - og) ** 2)
                               for dg, og in zip(dummy_gradients, observed_gradients)])

            # Update dummy data to minimize difference
            data_grad = compute_gradient_wrt_data(gradient_diff, dummy_data)
            dummy_data = optimizer.update(dummy_data, data_grad)

        return dummy_data


# Defense: Differential privacy noise prevents exact reconstruction

Membership Inference Attack

Determine if a specific sample was in the training set:

class MembershipInferenceAttack:
    """
    Infer whether a specific data point was in training set

    Based on model's confidence on that point
    """

    def __init__(self, shadow_models):
        """
        Args:
            shadow_models: Models trained on known datasets for attack training
        """
        self.shadow_models = shadow_models
        self.attack_model = self.train_attack_model()

    def train_attack_model(self):
        """
        Train meta-classifier to distinguish training vs non-training samples

        Uses prediction confidence as features
        """
        # Collect confidence scores from shadow models
        training_confidences = []
        non_training_confidences = []

        for shadow_model, train_data, test_data in self.shadow_models:
            # Confidence on training data (label=1: member)
            train_conf = shadow_model.predict_proba(train_data)
            training_confidences.extend(train_conf)

            # Confidence on non-training data (label=0: non-member)
            test_conf = shadow_model.predict_proba(test_data)
            non_training_confidences.extend(test_conf)

        # Train binary classifier
        X = np.vstack([training_confidences, non_training_confidences])
        y = np.array([1] * len(training_confidences) + [0] * len(non_training_confidences))

        attack_model = LogisticRegression()
        attack_model.fit(X, y)

        return attack_model

    def infer_membership(self, target_model, sample):
        """
        Predict if sample was in target model's training set

        Args:
            target_model: Model under attack
            sample: Sample to check

        Returns:
            Probability of membership
        """
        confidence = target_model.predict_proba(sample)
        membership_prob = self.attack_model.predict_proba([confidence])[0][1]

        return membership_prob


# Defense: Differential privacy reduces confidence, making inference harder

Sybil Attacks

Attack Description

Single adversary creates multiple fake identities to gain disproportionate influence:

class SybilAttack:
    """
    Sybil attack: Create multiple fake clients

    Adversary controls many identities to dominate aggregation
    """

    def __init__(self, num_sybils=10):
        self.num_sybils = num_sybils
        self.sybil_ids = [f"sybil_{i}" for i in range(num_sybils)]

    def generate_coordinated_attack(self, malicious_update):
        """
        All Sybil clients send the same malicious update

        With enough Sybils, can dominate averaging
        """
        return {sybil_id: malicious_update.copy()
                for sybil_id in self.sybil_ids}


# Example: 100 honest clients + 50 Sybils sending malicious update
# Malicious update gets 50/(100+50) = 33% weight in FedAvg

Defense: Client Validation

class ClientValidator:
    """
    Validate client identities to prevent Sybil attacks
    """

    def __init__(self):
        self.verified_clients = set()
        self.client_behaviors = {}

    def validate_client(self, client_id, credentials):
        """
        Verify client is legitimate

        Methods:
        - Device attestation (hardware signatures)
        - CAPTCHA / proof-of-work
        - Reputation systems
        - Rate limiting (one client per device)
        """

        # Check hardware attestation
        if not self.verify_attestation(credentials):
            return False

        # Check if already registered
        if client_id in self.verified_clients:
            return True

        # Perform additional verification (e.g., CAPTCHA)
        if self.perform_captcha_verification():
            self.verified_clients.add(client_id)
            return True

        return False

    def detect_sybil_behavior(self, client_updates):
        """
        Detect coordinated behavior indicating Sybils

        Sybils often send very similar updates
        """
        # Compute pairwise update similarities
        similarities = {}

        for i, (id1, update1) in enumerate(client_updates.items()):
            for id2, update2 in list(client_updates.items())[i+1:]:
                similarity = cosine_similarity(update1, update2)

                if similarity > 0.99:  # Very similar updates
                    if id1 not in similarities:
                        similarities[id1] = []
                    similarities[id1].append(id2)

        # Flag clients with many similar peers as potential Sybils
        suspected_sybils = [client for client, similar_to in similarities.items()
                           if len(similar_to) > 3]

        return suspected_sybils


def cosine_similarity(a, b):
    """Compute cosine similarity between two vectors"""
    return np.dot(a, b) / (np.linalg.norm(a) * np.linalg.norm(b) + 1e-10)

Comprehensive Defense Framework

class SecureFederatedLearningSystem:
    """
    Production-ready FL system with multiple defense layers
    """

    def __init__(self):
        # Defense components
        self.validator = ClientValidator()
        self.robust_aggregator = ByzantineRobustServer(
            aggregation_method='multi_krum',
            byzantine_ratio=0.2
        )
        self.privacy_mechanism = DifferentialPrivacy(epsilon=1.0, delta=1e-5)

        # Monitoring
        self.anomaly_detector = AnomalyDetector()
        self.attack_logger = AttackLogger()

    def training_round(self, client_updates):
        """
        Execute one secure training round

        Multi-layer defense:
        1. Client validation
        2. Anomaly detection
        3. Byzantine-robust aggregation
        4. Differential privacy
        """

        # Layer 1: Validate clients
        validated_updates = {}
        for client_id, update in client_updates.items():
            if self.validator.validate_client(client_id, update.get('credentials')):
                validated_updates[client_id] = update['model']

        # Layer 2: Detect Sybils
        suspected_sybils = self.validator.detect_sybil_behavior(validated_updates)
        for sybil in suspected_sybils:
            del validated_updates[sybil]
            self.attack_logger.log_sybil_attempt(sybil)

        # Layer 3: Anomaly detection
        updates_list = list(validated_updates.values())
        trusted_indices = self.anomaly_detector.detect_outliers(updates_list)
        trusted_updates = [updates_list[i] for i in trusted_indices]

        # Layer 4: Byzantine-robust aggregation
        if len(trusted_updates) > 0:
            global_update = self.robust_aggregator.robust_aggregate(trusted_updates)
        else:
            self.attack_logger.log_failed_round("No trusted updates")
            return None

        # Layer 5: Add differential privacy
        private_update = self.privacy_mechanism.add_noise(global_update)

        return private_update

    def monitor_and_adapt(self):
        """
        Monitor for attacks and adapt defenses

        Increase robustness if attacks detected
        """
        attack_rate = self.attack_logger.get_recent_attack_rate()

        if attack_rate > 0.1:  # More than 10% attack attempts
            # Strengthen defenses
            self.robust_aggregator.byzantine_ratio = min(0.4, self.robust_aggregator.byzantine_ratio + 0.1)
            self.anomaly_detector.threshold = max(2.0, self.anomaly_detector.threshold - 0.5)

            self.attack_logger.log_defense_adaptation(
                f"Increased robustness: attack_rate={attack_rate:.2%}"
            )

Best Practices for Secure FL

Threat Defense Trade-offs
Data Poisoning Anomaly detection, robust aggregation May reject honest outliers
Model Poisoning Byzantine-robust aggregation (Krum, median) Wastes some honest updates
Privacy Attacks Differential privacy, secure aggregation Reduced accuracy, higher overhead
Sybil Attacks Client validation, behavior analysis Deployment complexity
Inference Attacks DP, prediction perturbation Utility degradation

Chapter Summary

Review Questions

  1. Explain the difference between data poisoning and model poisoning attacks.
  2. How does a backdoor attack work? Why is it dangerous?
  3. Describe three types of Byzantine attacks on model aggregation.
  4. How can gradient inversion reconstruct training data? What defenses exist?
  5. What is a membership inference attack? How does differential privacy help?
  6. Explain Sybil attacks. How can client validation prevent them?
  7. Design a multi-layer defense system for a healthcare FL application.
  8. What trade-offs exist between security and model utility?
  9. How would you detect if your FL system is under attack in production?
  10. Why is Byzantine-robust aggregation necessary even with client validation?

Korea Standardization Infrastructure Mapping

Korea operates a comprehensive standards governance system through inter-ministerial cooperation. National Standards Council (under Prime Minister's Office, per Framework Act on National Standards Article 5) coordinates KATS (Korean Agency for Technology and Standards), MFDS (Ministry of Food and Drug Safety), MOTIE (Ministry of Trade, Industry and Energy), MSIT (Ministry of Science and ICT), MOIS (Ministry of the Interior and Safety), MOE (Ministry of Environment), MOHW (Ministry of Health and Welfare), MND (Ministry of National Defense), MCST (Ministry of Culture, Sports and Tourism), MOFA (Ministry of Foreign Affairs), MOJ (Ministry of Justice), and FSC (Financial Services Commission). Accreditation and Testing: KOLAS (Korea Laboratory Accreditation Scheme) accredits 800+ testing laboratories. KAS (Korea Accreditation System) accredits 50+ certification bodies. KTC (Korea Testing Certification), KTR (Korea Testing & Research Institute), KTL (Korea Testing Laboratory), and KCL (Korea Conformity Laboratories) provide conformance testing. Telecom and Cyber: KCC (Korea Communications Commission), KCA (Korea Communications Agency), TTA (Telecommunications Technology Association), IITP (Institute for Information & Communications Technology Planning & Evaluation), NIPA (National IT Industry Promotion Agency), KISA (Korea Internet & Security Agency), KCMVP (Korea Cryptographic Module Validation Program), NIS (National Intelligence Service), NSR (National Security Research Institute), and NCSC (National Cyber Security Center). National R&D Centers: KIST, ETRI, KAIST, Seoul National University, Yonsei University, Korea University, POSTECH, UNIST, GIST, DGIST, KISTI, KIER, KIMM, KRICT, KFRI, KRIBB. International Standards Cooperation: ISO TC/SC Korean secretariats, IEC TC/SC Korean secretariats, ITU-T Study Group Korean chairs, 3GPP RAN/SA Korean chairs, IEEE 802 Korean chairs, W3C Korea office, OASIS Korea office, IETF Korea cooperation, OECD CSTP, UN ESCAP, APEC SCSC Korean cooperation. Korean Industrial Standards (KS) Catalog: KS X (Information) 25,000+, KS A (Basic) 15,000+, KS B (Machinery) 25,000+, KS C (Electrical) 18,000+, KS D (Metallurgy) 12,000+, KS E (Mining) 5,000+, KS F (Construction) 18,000+, KS H (Food) 8,000+, KS I (Environment) 5,000+, KS J (Biology) 3,000+, KS K (Textile) 15,000+, KS L (Ceramics) 7,000+, KS M (Chemistry) 12,000+, KS P (Medical) 5,000+, KS Q (Quality Mgmt) 4,000+, KS R (Transport) 12,000+, KS S (Service) 3,000+, KS T (Packaging) 4,000+, KS V (Shipbuilding) 5,000+, KS W (Aerospace) 3,000+ โ€” totaling 220,000+ Korean Industrial Standards. Key Acts: Personal Information Protection Act (Act 19234, effective Sept 15, 2024), Electronic Government Act, Electronic Signature Act, Act on Promotion of Information and Communications Network Utilization and Information Protection, Information and Communications Infrastructure Protection Act, Data Industry Act, Public Data Act, AI Framework Act (Act 20212, effective July 2026), Industrial Technology Innovation Promotion Act, Framework Act on Science and Technology โ€” 70+ Korean standardization-related laws.