// Copyright 2022 The Chromium Authors
// Use of this source code is governed by a BSD-style license that can be
// found in the LICENSE file.

#include "components/safe_browsing/content/renderer/phishing_classifier/phishing_model_setter_impl.h"

#include "components/safe_browsing/core/common/phishing_classifier/scorer.h"
#include "third_party/blink/public/common/associated_interfaces/associated_interface_registry.h"

namespace safe_browsing {

PhishingModelSetterImpl::PhishingModelSetterImpl() = default;
PhishingModelSetterImpl::~PhishingModelSetterImpl() = default;

void PhishingModelSetterImpl::RegisterMojoInterfaces(
    blink::AssociatedInterfaceRegistry* associated_interfaces) {
  associated_interfaces->AddInterface<mojom::PhishingModelSetter>(
      base::BindRepeating(&PhishingModelSetterImpl::OnRendererAssociatedRequest,
                          base::Unretained(this)));
}

void PhishingModelSetterImpl::UnregisterMojoInterfaces(
    blink::AssociatedInterfaceRegistry* associated_interfaces) {
  associated_interfaces->RemoveInterface(mojom::PhishingModelSetter::Name_);
}

void PhishingModelSetterImpl::SetImageEmbeddingAndPhishingTfLiteModel(
    int classification_input_width,
    int classification_input_height,
    base::File classification_model,
    int image_embedding_input_width,
    int image_embedding_input_height,
    base::File image_embedding_model) {
  std::unique_ptr<Scorer> scorer =
      safe_browsing::Scorer::CreateScorerWithImageEmbeddingModel(
          classification_input_width, classification_input_height,
          std::move(classification_model), image_embedding_input_width,
          image_embedding_input_height, std::move(image_embedding_model));

  if (!scorer) {
    return;
  }

  ScorerStorage::GetInstance()->SetScorer(std::move(scorer));

  if (observer_for_testing_.is_bound()) {
    observer_for_testing_->PhishingModelUpdated();
  }
}

void PhishingModelSetterImpl::SetPhishingTfLiteModel(
    int classification_input_width,
    int classification_input_height,
    base::File tflite_visual_model) {
  std::unique_ptr<Scorer> scorer = safe_browsing::Scorer::Create(
      classification_input_width, classification_input_height,
      std::move(tflite_visual_model));

  if (!scorer) {
    return;
  }

  ScorerStorage::GetInstance()->SetScorer(std::move(scorer));

  if (observer_for_testing_.is_bound()) {
    observer_for_testing_->PhishingModelUpdated();
  }
}

void PhishingModelSetterImpl::AttachImageEmbeddingModelAndDimensions(
    int image_embedding_input_width,
    int image_embedding_input_height,
    base::File image_embedding_model) {
  Scorer* scorer = ScorerStorage::GetInstance()->GetScorer();
  if (!scorer) {
    return;
  }

  scorer->AttachImageEmbeddingModel(image_embedding_input_width,
                                    image_embedding_input_height,
                                    std::move(image_embedding_model));
}

void PhishingModelSetterImpl::ClearScorer() {
  ScorerStorage::GetInstance()->ClearScorer();
}

void PhishingModelSetterImpl::SetTestObserver(
    mojo::PendingRemote<mojom::PhishingModelSetterTestObserver> observer,
    SetTestObserverCallback callback) {
  if (observer_for_testing_.is_bound())
    observer_for_testing_.reset();
  observer_for_testing_.Bind(std::move(observer));
  std::move(callback).Run();
}

void PhishingModelSetterImpl::OnRendererAssociatedRequest(
    mojo::PendingAssociatedReceiver<mojom::PhishingModelSetter> receiver) {
  receiver_.reset();
  receiver_.Bind(std::move(receiver));
}

}  // namespace safe_browsing
