#!/usr/bin/env perl

use strict;
use warnings;
use FindBin qw($Bin);
use lib "$Bin/../lib";
use Getopt::Long qw(GetOptions);
use Storable qw(retrieve nfreeze thaw);
use LaptopNN::Network;

binmode STDIN,  ':raw';
binmode STDOUT, ':raw';
binmode STDERR, ':encoding(UTF-8)';

my $state_path;
GetOptions('state=s' => \$state_path)
    or die "invalid worker arguments\n";
die "--state is required\n" unless defined $state_path;

my $state = retrieve($state_path);
die "invalid worker state\n"
    unless ref($state) eq 'HASH'
        && ref($state->{X}) eq 'ARRAY'
        && ref($state->{y}) eq 'ARRAY'
        && ref($state->{network}) eq 'HASH';
my $network = LaptopNN::Network->from_hash($state->{network});

while (1) {
    my $payload = _read_frame(*STDIN);
    last unless defined $payload;
    my $request = eval { thaw($payload) };
    if ($@ || ref($request) ne 'HASH') {
        _write_frame(*STDOUT, nfreeze({ ok => 0, error => "invalid request: $@" }));
        next;
    }
    if (($request->{command} // '') eq 'stop') {
        _write_frame(*STDOUT, nfreeze({ ok => 1, stopped => 1 }));
        last;
    }

    my $response = eval {
        die "unsupported worker command\n"
            unless ($request->{command} // '') eq 'gradient';
        die "worker request has no rows\n"
            unless ref($request->{rows}) eq 'ARRAY' && @{ $request->{rows} };
        $network->_unpack_parameters($request->{weights});
        srand($request->{seed} // 42);
        my $gradient = $network->_zero_gradients;
        my $weight_sum = 0;
        for my $row (@{ $request->{rows} }) {
            die "worker row index is out of range\n"
                if $row < 0 || $row > $#{ $state->{X} };
            my $target = $state->{y}[$row];
            my $weight = defined($request->{class_weights})
                ? $request->{class_weights}[$target]
                : 1;
            $network->_accumulate_sample_gradient(
                $state->{X}[$row],
                $target,
                $gradient,
                sample_weight   => $weight,
                label_smoothing => $request->{label_smoothing},
                training        => 1,
            );
            $weight_sum += $weight;
        }
        {
            ok          => 1,
            task_number => $request->{task_number},
            gradients   => $network->_pack_gradients($gradient),
            weight_sum  => $weight_sum,
        };
    };
    if ($@) {
        $response = {
            ok          => 0,
            task_number => $request->{task_number},
            error       => "$@",
        };
    }
    _write_frame(*STDOUT, nfreeze($response));
}

sub _write_frame {
    my ($fh, $payload) = @_;
    local $SIG{PIPE} = 'IGNORE';
    my $frame = pack('N', length($payload)) . $payload;
    my $offset = 0;
    while ($offset < length($frame)) {
        my $written = syswrite($fh, $frame, length($frame) - $offset, $offset);
        die "worker cannot write response: $!\n" unless defined($written) && $written > 0;
        $offset += $written;
    }
}

sub _read_frame {
    my ($fh) = @_;
    my $header = _read_exact($fh, 4, 1);
    return undef unless defined $header;
    my $length = unpack('N', $header);
    die "invalid request frame length\n" if $length < 1 || $length > 1_000_000_000;
    return _read_exact($fh, $length, 0);
}

sub _read_exact {
    my ($fh, $wanted, $allow_clean_eof) = @_;
    my $buffer = '';
    while (length($buffer) < $wanted) {
        my $read = sysread($fh, $buffer, $wanted - length($buffer), length($buffer));
        die "worker cannot read request: $!\n" unless defined $read;
        return undef if $read == 0 && $allow_clean_eof && length($buffer) == 0;
        die "worker received a truncated request\n" if $read == 0;
    }
    return $buffer;
}
