diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS index 7fba5cd..2d84450 100644 --- a/.github/CODEOWNERS +++ b/.github/CODEOWNERS @@ -1,2 +1,2 @@ # This should match the MAINTAINERS.md file in the root of this repo -* @msfroh @macohen @noCharger @mingshl \ No newline at end of file +* @msfroh @macohen @noCharger @mingshl @sejli \ No newline at end of file diff --git a/.github/workflows/CI.yml b/.github/workflows/CI.yml index 61bfc57..6a31d4d 100644 --- a/.github/workflows/CI.yml +++ b/.github/workflows/CI.yml @@ -15,7 +15,7 @@ jobs: Build-search-request-processor: strategy: matrix: - java: [11, 17] + java: [21] os: [ubuntu-latest, macos-latest, windows-latest] name: Build and Test Search Request Processor Plugin @@ -23,10 +23,10 @@ jobs: steps: - name: Checkout Search Request Processor - uses: actions/checkout@v1 + uses: actions/checkout@50fbc622fc4ef5163becd7fab6573eac35f8462e # v1 - name: Setup Java ${{ matrix.java }} - uses: actions/setup-java@v1 + uses: actions/setup-java@b6e674f4b717d7b0ae3baee0fbe79f498905dfde # v1 with: java-version: ${{ matrix.java }} @@ -42,7 +42,7 @@ jobs: - name: Upload Coverage Report if: ${{matrix.os}} == 'ubuntu' - uses: codecov/codecov-action@v1 + uses: codecov/codecov-action@29386c70ef20e286228c72b668a06fd0e8399192 # v1 with: token: ${{ secrets.CODECOV_TOKEN }} diff --git a/.github/workflows/add-untriaged.yml b/.github/workflows/add-untriaged.yml index 15b9a55..846441e 100644 --- a/.github/workflows/add-untriaged.yml +++ b/.github/workflows/add-untriaged.yml @@ -4,11 +4,14 @@ on: issues: types: [opened, reopened, transferred] +permissions: + issues: write + jobs: apply-label: runs-on: ubuntu-latest steps: - - uses: actions/github-script@v6 + - uses: actions/github-script@3a2844b7e9c422d3c10d287c895573f7108da1b3 # v9 with: script: | github.rest.issues.addLabels({ diff --git a/.github/workflows/backport.yml b/.github/workflows/backport.yml index e47d8d8..2609a3c 100644 --- a/.github/workflows/backport.yml +++ b/.github/workflows/backport.yml @@ -15,14 +15,14 @@ jobs: steps: - name: GitHub App token id: github_app_token - uses: tibdex/github-app-token@v1.5.0 + uses: tibdex/github-app-token@1901dc7d52169e70c27a8da37aef0d423e2867a2 # v1.5.0 with: app_id: ${{ secrets.APP_ID }} private_key: ${{ secrets.APP_PRIVATE_KEY }} installation_id: 22958780 - name: Backport - uses: VachaShah/backport@v1.1.4 + uses: VachaShah/backport@28c49d91ceec57d7c9f625f1031c1a4d637251f5 # v1.1.4 with: github_token: ${{ steps.github_app_token.outputs.token }} branch_name: backport/backport-${{ github.event.number }} diff --git a/.github/workflows/backwards_compatibility_tests_workflow.yml b/.github/workflows/backwards_compatibility_tests_workflow.yml index 6e5217c..8644c7c 100644 --- a/.github/workflows/backwards_compatibility_tests_workflow.yml +++ b/.github/workflows/backwards_compatibility_tests_workflow.yml @@ -13,8 +13,8 @@ jobs: Restart-Upgrade-BWCTests-k-NN: strategy: matrix: - java: [ 11, 17 ] - bwc_version : [ "2.6.0" ] + java: [ 21 ] + bwc_version : [ "2.7.0" ] opensearch_version : [ "3.0.0-SNAPSHOT" ] name: SRP Restart-Upgrade BWC Tests @@ -24,10 +24,10 @@ jobs: steps: - name: Checkout SRP - uses: actions/checkout@v1 + uses: actions/checkout@50fbc622fc4ef5163becd7fab6573eac35f8462e # v1 - name: Setup Java ${{ matrix.java }} - uses: actions/setup-java@v1 + uses: actions/setup-java@b6e674f4b717d7b0ae3baee0fbe79f498905dfde # v1 with: java-version: ${{ matrix.java }} @@ -41,8 +41,8 @@ jobs: Rolling-Upgrade-BWCTests-SRP: strategy: matrix: - java: [ 11, 17 ] - bwc_version: [ "2.6.0" ] + java: [ 21 ] + bwc_version: [ "2.11.1" ] opensearch_version: [ "3.0.0-SNAPSHOT" ] name: SRP Rolling-Upgrade BWC Tests @@ -52,10 +52,10 @@ jobs: steps: - name: Checkout SRP - uses: actions/checkout@v1 + uses: actions/checkout@50fbc622fc4ef5163becd7fab6573eac35f8462e # v1 - name: Setup Java ${{ matrix.java }} - uses: actions/setup-java@v1 + uses: actions/setup-java@b6e674f4b717d7b0ae3baee0fbe79f498905dfde # v1 with: java-version: ${{ matrix.java }} diff --git a/.github/workflows/create-documentation-issue.yml b/.github/workflows/create-documentation-issue.yml index cb1eb40..3938005 100644 --- a/.github/workflows/create-documentation-issue.yml +++ b/.github/workflows/create-documentation-issue.yml @@ -14,14 +14,14 @@ jobs: steps: - name: GitHub App token id: github_app_token - uses: tibdex/github-app-token@v1.5.0 + uses: tibdex/github-app-token@1901dc7d52169e70c27a8da37aef0d423e2867a2 # v1.5.0 with: app_id: ${{ secrets.APP_ID }} private_key: ${{ secrets.APP_PRIVATE_KEY }} installation_id: 22958780 - name: Checkout code - uses: actions/checkout@v2 + uses: actions/checkout@ee0669bd1cc54295c223e0bb666b733df41de1c5 # v2 - name: Edit the issue template run: | @@ -29,7 +29,7 @@ jobs: - name: Create Issue From File id: create-issue - uses: peter-evans/create-issue-from-file@v4 + uses: peter-evans/create-issue-from-file@433e51abf769039ee20ba1293a088ca19d573b7f # v4 with: title: Add documentation related to new feature content-filepath: ./.github/ISSUE_TEMPLATE/documentation-issue.md diff --git a/.github/workflows/delete_backport_branch.yml b/.github/workflows/delete_backport_branch.yml index a97f9cd..156d749 100644 --- a/.github/workflows/delete_backport_branch.yml +++ b/.github/workflows/delete_backport_branch.yml @@ -10,6 +10,6 @@ jobs: if: startsWith(github.event.pull_request.head.ref,'backport-') steps: - name: Delete merged branch - uses: SvanBoxel/delete-merged-branch@main + uses: SvanBoxel/delete-merged-branch@2b5b058e3db41a3328fd9a6a58fd4c2545a14353 # main env: GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} \ No newline at end of file diff --git a/.github/workflows/dependabot_pr.yml b/.github/workflows/dependabot_pr.yml index bdafb18..7d13551 100644 --- a/.github/workflows/dependabot_pr.yml +++ b/.github/workflows/dependabot_pr.yml @@ -11,14 +11,14 @@ jobs: steps: - name: GitHub App token id: github_app_token - uses: tibdex/github-app-token@v1.5.0 + uses: tibdex/github-app-token@1901dc7d52169e70c27a8da37aef0d423e2867a2 # v1.5.0 with: app_id: ${{ secrets.APP_ID }} private_key: ${{ secrets.APP_PRIVATE_KEY }} installation_id: 22958780 - name: Check out code - uses: actions/checkout@v2 + uses: actions/checkout@ee0669bd1cc54295c223e0bb666b733df41de1c5 # v2 with: token: ${{ steps.github_app_token.outputs.token }} @@ -27,7 +27,7 @@ jobs: ./gradlew updateSHAs - name: Commit the changes - uses: stefanzweifel/git-auto-commit-action@v4.7.2 + uses: stefanzweifel/git-auto-commit-action@3ea6ae190baf489ba007f7c92608f33ce20ef04a # v4.7.2 with: commit_message: Updating SHAs branch: ${{ github.head_ref }} @@ -36,7 +36,7 @@ jobs: commit_options: '--signoff' - name: Commit the changes - uses: stefanzweifel/git-auto-commit-action@v4.7.2 + uses: stefanzweifel/git-auto-commit-action@3ea6ae190baf489ba007f7c92608f33ce20ef04a # v4.7.2 with: commit_message: Spotless formatting branch: ${{ github.head_ref }} @@ -45,12 +45,12 @@ jobs: commit_options: '--signoff' - name: Update the changelog - uses: dangoslen/dependabot-changelog-helper@v1 + uses: dangoslen/dependabot-changelog-helper@780f7c82213ff956b1bd8cb484ba67d1fbe8b4ba # v1 with: version: 'Unreleased' - name: Commit the changes - uses: stefanzweifel/git-auto-commit-action@v4 + uses: stefanzweifel/git-auto-commit-action@3ea6ae190baf489ba007f7c92608f33ce20ef04a # v4 with: commit_message: "Update changelog" branch: ${{ github.head_ref }} diff --git a/.github/workflows/draft-release-notes-workflow.yml b/.github/workflows/draft-release-notes-workflow.yml index 6b3d89c..48e2bf9 100644 --- a/.github/workflows/draft-release-notes-workflow.yml +++ b/.github/workflows/draft-release-notes-workflow.yml @@ -11,7 +11,7 @@ jobs: runs-on: ubuntu-latest steps: - name: Update draft release notes - uses: release-drafter/release-drafter@v5 + uses: release-drafter/release-drafter@09c613e259eb8d4e7c81c2cb00618eb5fc4575a7 # v5 with: config-name: draft-release-notes-config.yml name: Version (set here) diff --git a/.github/workflows/links.yml b/.github/workflows/links.yml index 3d0b81a..88f0d32 100644 --- a/.github/workflows/links.yml +++ b/.github/workflows/links.yml @@ -11,10 +11,10 @@ jobs: runs-on: ubuntu-latest steps: - - uses: actions/checkout@v2 + - uses: actions/checkout@ee0669bd1cc54295c223e0bb666b733df41de1c5 # v2 - name: lychee Link Checker id: lychee - uses: lycheeverse/lychee-action@master + uses: lycheeverse/lychee-action@6da1d14f3a43098a294b7696d93d938aa8d20fc0 # master with: args: --accept=200,403,429 **/*.html **/*.md **/*.txt **/*.json env: diff --git a/.github/workflows/pr_stats.yml b/.github/workflows/pr_stats.yml index 96c971b..5cb0564 100644 --- a/.github/workflows/pr_stats.yml +++ b/.github/workflows/pr_stats.yml @@ -12,4 +12,4 @@ jobs: pull-requests: write steps: - name: Run pull request stats - uses: flowwer-dev/pull-request-stats@master + uses: flowwer-dev/pull-request-stats@0dde6edf8b7db75684533021212c6e85fb987dbb # master diff --git a/MAINTAINERS.md b/MAINTAINERS.md index db94666..12c0cd1 100644 --- a/MAINTAINERS.md +++ b/MAINTAINERS.md @@ -10,3 +10,4 @@ This document contains a list of maintainers in this repo. See [opensearch-proje | Mark Cohen | [macohen](https://github.com/macohen) | Amazon | | Louis Chu | [noCharger](https://github.com/noCharger) | Amazon | | Mingshi Liu | [mingshl](https://github.com/mingshl) | Amazon | +| Sean Li | [sejli](https://github.com/sejli) | Amazon | diff --git a/README.md b/README.md index 5842249..06b5ecc 100644 --- a/README.md +++ b/README.md @@ -2,7 +2,7 @@ [![codecov](https://codecov.io/gh/opensearch-project/search-processor/branch/main/graph/badge.svg?token=PYQO2GW39S)](https://codecov.io/gh/opensearch-project/search-processor) ![PRs welcome!](https://img.shields.io/badge/PRs-welcome!-success) -# Search Query & Request Transformers +# Search Rerankers: AWS Kendra & AWS Personalize - [Welcome!](#welcome) - [Project Resources](#project-resources) - [Code of Conduct](#code-of-conduct) @@ -10,21 +10,21 @@ - [Copyright](#copyright) ## Welcome! -This repository is the home of an evolving project that aims to create a pipeline of transformers to preprocess queries before search and post-process results after search. The first component here is a plugin to re-rank search results before returning them to the client inline. In the coming year, we will add hooks to configure other re-rankers and allow users to add their own components to the pipeline. Logging will also be a critical part of the pipeline in two ways: -1. Logging information about the search experience (e.g. query, search results returned from the index, search results returned to the OpenSearch client) -1. Logging debug information about the transformers +This repository hosts the code for two self-install re-rankers that integrate into [Search Pipelines](https://opensearch.org/docs/latest/search-plugins/search-pipelines/index/). User documentation for the Personalize Reranker is [here](https://opensearch.org/docs/latest/search-plugins/search-pipelines/personalize-search-ranking/). For Kendra, it is [here](https://opensearch.org/docs/latest/search-plugins/search-relevance/index/#reranking-results-with-kendra-intelligent-ranking-for-opensearch). + +# Search Processors: Where Do They Go? +The current guideline for developing processors is that if you are developing a processor that would introduce new dependencies in [OpenSearch Core](https://github.com/opensearch-project/OpenSearch) (e.g. new libraries, makes a network connection outside of OpenSearch), it should be in a separate repository. Please consider creating it in a standalone repository since each processor should be thought of like a \*NIX command with input and output connected by pipes (i.e. a Search Pipeline). Each processor should do one thing and do it well. Otherwise, it could go into the OpenSearch repository under [org.opensearch.search.pipeline.common](https://github.com/opensearch-project/OpenSearch/tree/a08d588691c3b232e65d73b0a0c2fc5c72c870cf/modules/search-pipeline-common). If you have doubts, just create an issue in OpenSearch Core and, if you have one, a new PR. Maintainers will help guide you. -We will be publishing an RFC soon to give more detail and have a deeper conversation, but for now take a look at the code, open issues, comment, etc. # History This repository has also been used for discussion and ideas around search relevance. These discussions still exist here, however due to the relatively new standard of having one repo per plugin in OpenSearch and our implementations beginning to make it into the OpenSearch build, we have two repositories now. This repository will develop into a plugin that will allow OpenSearch users to rewrite search queries, rerank results, and log data about those actions. The other repository, [dashboards-search-relevance](https://www.github.com/opensearch-projects/dashboards-search-relevance), is where we will build front-end tooling to help relevance engineers and business users tune results. - ## Project Resources * [OpenSearch Project Website](https://opensearch.org/) * [Downloads](https://opensearch.org/downloads.html) * [Project Principles](https://opensearch.org/#principles) +* [Search Pipelines](https://opensearch.org/docs/latest/search-plugins/search-pipelines/index/) * [Contributing to OpenSearch Search Request Processor](CONTRIBUTING.md) * [Search Relevance](RELEVANCE.md) * [Maintainer Responsibilities](MAINTAINERS.md) diff --git a/amazon-kendra-intelligent-ranking/build.gradle b/amazon-kendra-intelligent-ranking/build.gradle new file mode 100644 index 0000000..60821f6 --- /dev/null +++ b/amazon-kendra-intelligent-ranking/build.gradle @@ -0,0 +1,118 @@ +import org.opensearch.gradle.test.RestIntegTestTask + +apply plugin: 'java' +apply plugin: 'idea' +apply plugin: 'opensearch.opensearchplugin' +apply plugin: 'opensearch.yaml-rest-test' +apply plugin: 'jacoco' + +group = 'org.opensearch' + +def pluginName = 'amazon-kendra-intelligent-ranking' +def pluginDescription = 'Rerank search results using Amazon Kendra Intelligent Ranking' +def projectPath = 'org.opensearch' +def pathToPlugin = 'search.relevance' +def pluginClassName = 'AmazonKendraIntelligentRankingPlugin' + +opensearchplugin { + name "opensearch-${pluginName}-${plugin_version}.0" + version "${plugin_version}" + description pluginDescription + classname "${projectPath}.${pathToPlugin}.${pluginClassName}" + licenseFile rootProject.file('LICENSE') + noticeFile rootProject.file('NOTICE') +} + +java { + targetCompatibility = JavaVersion.VERSION_21 + sourceCompatibility = JavaVersion.VERSION_21 +} + +// This requires an additional Jar not published as part of build-tools +loggerUsageCheck.enabled = false + +// No need to validate pom, as we do not upload to maven/sonatype +validateNebulaPom.enabled = false + +buildscript { + repositories { + mavenLocal() + maven { url "https://aws.oss.sonatype.org/content/repositories/snapshots" } + mavenCentral() + maven { url "https://plugins.gradle.org/m2/" } + } + + dependencies { + classpath "org.opensearch.gradle:build-tools:${opensearch_version}" + } +} + +repositories { + mavenLocal() + maven { url "https://aws.oss.sonatype.org/content/repositories/snapshots" } + mavenCentral() + maven { url "https://plugins.gradle.org/m2/" } +} + +dependencies { + implementation 'com.ibm.icu:icu4j:57.2' + implementation 'org.apache.httpcomponents:httpclient:4.5.14' + implementation 'org.apache.httpcomponents:httpcore:4.4.16' + implementation 'com.fasterxml.jackson.core:jackson-databind:2.18.2' + implementation 'com.fasterxml.jackson.core:jackson-core:2.18.2' + implementation 'com.fasterxml.jackson.core:jackson-annotations:2.18.2' + implementation 'commons-logging:commons-logging:1.2' + implementation 'com.amazonaws:aws-java-sdk-sts:1.12.300' + implementation 'com.amazonaws:aws-java-sdk-core:1.12.300' +} + + +allprojects { + plugins.withId('jacoco') { + jacoco.toolVersion = '0.8.9' + } +} + + +test { + include '**/*Tests.class' + finalizedBy jacocoTestReport +} + +task integTest(type: RestIntegTestTask) { + description = "Run tests against a cluster" + testClassesDirs = sourceSets.test.output.classesDirs + classpath = sourceSets.test.runtimeClasspath +} +tasks.named("check").configure { dependsOn(integTest) } + +integTest { + // The --debug-jvm command-line option makes the cluster debuggable; this makes the tests debuggable + if (System.getProperty("test.debug") != null) { + jvmArgs '-agentlib:jdwp=transport=dt_socket,server=y,suspend=y,address=*:5005' + } +} + +testClusters.integTest { + testDistribution = "ARCHIVE" + + // This installs our plugin into the testClusters + plugin(project.tasks.bundlePlugin.archiveFile) +} + +run { + useCluster testClusters.integTest +} + +jacocoTestReport { + dependsOn test + reports { + xml.required = true + html.required = true + } +} + +// TODO: Enable these checks +dependencyLicenses.enabled = false +thirdPartyAudit.enabled = false +loggerUsageCheck.enabled = false diff --git a/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/AmazonKendraIntelligentRankingPlugin.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/AmazonKendraIntelligentRankingPlugin.java new file mode 100644 index 0000000..ef9df5d --- /dev/null +++ b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/AmazonKendraIntelligentRankingPlugin.java @@ -0,0 +1,119 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance; + +import org.opensearch.action.support.ActionFilter; +import org.opensearch.client.Client; +import org.opensearch.cluster.metadata.IndexNameExpressionResolver; +import org.opensearch.cluster.service.ClusterService; +import org.opensearch.common.settings.Setting; +import org.opensearch.core.common.io.stream.NamedWriteableRegistry; +import org.opensearch.core.xcontent.NamedXContentRegistry; +import org.opensearch.env.Environment; +import org.opensearch.env.NodeEnvironment; +import org.opensearch.plugins.ActionPlugin; +import org.opensearch.plugins.Plugin; +import org.opensearch.plugins.SearchPipelinePlugin; +import org.opensearch.plugins.SearchPlugin; +import org.opensearch.repositories.RepositoriesService; +import org.opensearch.script.ScriptService; +import org.opensearch.search.pipeline.Processor; +import org.opensearch.search.pipeline.SearchResponseProcessor; +import org.opensearch.search.relevance.actionfilter.SearchActionFilter; +import org.opensearch.search.relevance.client.OpenSearchClient; +import org.opensearch.search.relevance.configuration.ResultTransformerConfigurationFactory; +import org.opensearch.search.relevance.configuration.SearchConfigurationExtBuilder; +import org.opensearch.search.relevance.transformer.ResultTransformer; +import org.opensearch.search.relevance.transformer.kendraintelligentranking.KendraIntelligentRanker; +import org.opensearch.search.relevance.transformer.kendraintelligentranking.client.KendraClientSettings; +import org.opensearch.search.relevance.transformer.kendraintelligentranking.client.KendraHttpClient; +import org.opensearch.search.relevance.transformer.kendraintelligentranking.configuration.KendraIntelligentRankerSettings; +import org.opensearch.search.relevance.transformer.kendraintelligentranking.configuration.KendraIntelligentRankingConfigurationFactory; +import org.opensearch.search.relevance.transformer.kendraintelligentranking.pipeline.KendraRankingResponseProcessor; +import org.opensearch.threadpool.ThreadPool; +import org.opensearch.watcher.ResourceWatcherService; + +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collection; +import java.util.List; +import java.util.Map; +import java.util.function.Supplier; +import java.util.stream.Collectors; + +public class AmazonKendraIntelligentRankingPlugin extends Plugin implements ActionPlugin, SearchPlugin, SearchPipelinePlugin { + + private OpenSearchClient openSearchClient; + private KendraHttpClient kendraClient; + private KendraIntelligentRanker kendraIntelligentRanker; + private KendraClientSettings kendraClientSettings; + + private Collection getAllResultTransformers() { + // Initialize and add other transformers here + return List.of(this.kendraIntelligentRanker); + } + + private Collection getResultTransformerConfigurationFactories() { + return List.of(KendraIntelligentRankingConfigurationFactory.INSTANCE); + } + + @Override + public List getActionFilters() { + return Arrays.asList(new SearchActionFilter(getAllResultTransformers(), openSearchClient)); + } + + @Override + public List> getSettings() { + // NOTE: cannot use kendraIntelligentRanker.getTransformerSettings because the object is not yet created + List> allTransformerSettings = new ArrayList<>(); + allTransformerSettings.addAll(KendraIntelligentRankerSettings.getAllSettings()); + // Add settings for other transformers here + return allTransformerSettings; + } + + @Override + public Collection createComponents( + Client client, + ClusterService clusterService, + ThreadPool threadPool, + ResourceWatcherService resourceWatcherService, + ScriptService scriptService, + NamedXContentRegistry xContentRegistry, + Environment environment, + NodeEnvironment nodeEnvironment, + NamedWriteableRegistry namedWriteableRegistry, + IndexNameExpressionResolver indexNameExpressionResolver, + Supplier repositoriesServiceSupplier + ) { + this.openSearchClient = new OpenSearchClient(client); + this.kendraClientSettings = KendraClientSettings.getClientSettings(environment.settings()); + this.kendraClient = new KendraHttpClient(this.kendraClientSettings); + this.kendraIntelligentRanker = new KendraIntelligentRanker(this.kendraClient); + + return Arrays.asList( + this.openSearchClient, + this.kendraClientSettings, + this.kendraClient, + this.kendraIntelligentRanker + ); + } + + @Override + public List> getSearchExts() { + Map resultTransformerMap = getResultTransformerConfigurationFactories().stream() + .collect(Collectors.toMap(ResultTransformerConfigurationFactory::getName, i -> i)); + return List.of(new SearchExtSpec<>(SearchConfigurationExtBuilder.NAME, + input -> new SearchConfigurationExtBuilder(input, resultTransformerMap), + parser -> SearchConfigurationExtBuilder.parse(parser, resultTransformerMap))); + } + + @Override + public Map> getResponseProcessors(Parameters parameters) { + return Map.of(KendraRankingResponseProcessor.TYPE, new KendraRankingResponseProcessor.Factory(this.kendraClientSettings)); + } +} \ No newline at end of file diff --git a/src/main/java/org/opensearch/search/relevance/actionfilter/SearchActionFilter.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/actionfilter/SearchActionFilter.java similarity index 81% rename from src/main/java/org/opensearch/search/relevance/actionfilter/SearchActionFilter.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/actionfilter/SearchActionFilter.java index 85bfb06..229172c 100644 --- a/src/main/java/org/opensearch/search/relevance/actionfilter/SearchActionFilter.java +++ b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/actionfilter/SearchActionFilter.java @@ -10,9 +10,7 @@ import org.apache.logging.log4j.LogManager; import org.apache.logging.log4j.Logger; import org.opensearch.OpenSearchException; -import org.opensearch.action.ActionListener; import org.opensearch.action.ActionRequest; -import org.opensearch.action.ActionResponse; import org.opensearch.action.search.SearchAction; import org.opensearch.action.search.SearchRequest; import org.opensearch.action.search.SearchResponse; @@ -20,10 +18,13 @@ import org.opensearch.action.support.ActionFilter; import org.opensearch.action.support.ActionFilterChain; import org.opensearch.common.io.stream.BytesStreamOutput; -import org.opensearch.common.io.stream.NamedWriteableAwareStreamInput; -import org.opensearch.common.io.stream.NamedWriteableRegistry; -import org.opensearch.common.io.stream.StreamInput; +import org.opensearch.core.action.ActionListener; +import org.opensearch.core.action.ActionResponse; +import org.opensearch.core.common.io.stream.NamedWriteableAwareStreamInput; +import org.opensearch.core.common.io.stream.NamedWriteableRegistry; +import org.opensearch.core.common.io.stream.StreamInput; import org.opensearch.common.settings.Setting; +import org.opensearch.common.settings.Settings; import org.opensearch.search.SearchHit; import org.opensearch.search.SearchHits; import org.opensearch.search.aggregations.InternalAggregations; @@ -89,7 +90,7 @@ public void app // TODO: Remove originalSearchSource and replace with a deep copy of the SearchRequest object // once https://github.com/opensearch-project/OpenSearch/issues/869 is implemented - SearchSourceBuilder originalSearchSource = null; + SearchSourceBuilder originalSearchSource; if (searchRequest.source() != null) { originalSearchSource = searchRequest.source().shallowCopy(); if (searchRequest.source().fetchSource() != null) { @@ -98,6 +99,8 @@ public void app originalSearchSource.fetchSource(new FetchSourceContext(fetchSourceContext.fetchSource(), fetchSourceContext.includes(), fetchSourceContext.excludes())); } + } else { + originalSearchSource = null; } final String[] indices = searchRequest.indices(); @@ -107,27 +110,29 @@ public void app return; } - List resultTransformerConfigurations = - getResultTransformerConfigurations(indices[0], searchRequest); - LinkedHashMap orderedTransformersAndConfigs = new LinkedHashMap<>(); - for (ResultTransformerConfiguration config : resultTransformerConfigurations) { - ResultTransformer resultTransformer = resultTransformerMap.get(config.getTransformerName()); - // TODO: Should transformers make a decision based on the original request or the request they receive in the chain - if (resultTransformer.shouldTransform(searchRequest, config)) { - searchRequest = resultTransformer.preprocessRequest(searchRequest, config); - orderedTransformersAndConfigs.put(resultTransformer, config); + ActionListener> resultTransformerConfigsListener = ActionListener.wrap(rtc -> { + LinkedHashMap orderedTransformersAndConfigs = new LinkedHashMap<>(); + SearchRequest transformedRequest = searchRequest; + for (ResultTransformerConfiguration config : rtc) { + ResultTransformer resultTransformer = resultTransformerMap.get(config.getTransformerName()); + // TODO: Should transformers make a decision based on the original request or the request they receive in the chain + if (resultTransformer.shouldTransform(searchRequest, config)) { + transformedRequest = resultTransformer.preprocessRequest(transformedRequest, config); + orderedTransformersAndConfigs.put(resultTransformer, config); + } } - } - if (!orderedTransformersAndConfigs.isEmpty()) { - final ActionListener searchResponseListener = createSearchResponseListener( - listener, startTime, orderedTransformersAndConfigs, searchRequest, originalSearchSource); - chain.proceed(task, action, request, searchResponseListener); - return; - } + if (!orderedTransformersAndConfigs.isEmpty()) { + final ActionListener searchResponseListener = createSearchResponseListener( + listener, startTime, orderedTransformersAndConfigs, transformedRequest, originalSearchSource); + chain.proceed(task, action, request, searchResponseListener); + return; + } + chain.proceed(task, action, request, listener); + }, listener::onFailure); - chain.proceed(task, action, request, listener); + getResultTransformerConfigurations(indices[0], searchRequest, resultTransformerConfigsListener); } /** @@ -139,16 +144,18 @@ public void app * @return ordered and validated list of result transformers, empty list if not specified at * either request or index level */ - private List getResultTransformerConfigurations( + private void getResultTransformerConfigurations( final String indexName, - final SearchRequest searchRequest) { + final SearchRequest searchRequest, + ActionListener> resultTransformerConfigListener) { List configs = new ArrayList<>(); // Request level configuration takes precedence over index level configs = ConfigurationUtils.getResultTransformersFromRequestConfiguration(searchRequest); if (!configs.isEmpty()) { - return configs; + resultTransformerConfigListener.onResponse(configs); + return; } // Fetch all index settings for this plugin @@ -159,10 +166,10 @@ private List getResultTransformerConfigurations( .map(Setting::getKey)) .toArray(String[]::new); - configs = ConfigurationUtils.getResultTransformersFromIndexConfiguration( - openSearchClient.getIndexSettings(indexName, settingNames), resultTransformerMap); - return configs; + ActionListener settingsListener = ActionListener.map(resultTransformerConfigListener, + s -> ConfigurationUtils.getResultTransformersFromIndexConfiguration(s, resultTransformerMap)); + openSearchClient.getIndexSettings(indexName, settingNames, settingsListener); } /** diff --git a/src/main/java/org/opensearch/search/relevance/client/OpenSearchClient.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/client/OpenSearchClient.java similarity index 70% rename from src/main/java/org/opensearch/search/relevance/client/OpenSearchClient.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/client/OpenSearchClient.java index 1b306dc..42baf89 100644 --- a/src/main/java/org/opensearch/search/relevance/client/OpenSearchClient.java +++ b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/client/OpenSearchClient.java @@ -12,6 +12,7 @@ import org.opensearch.action.admin.indices.settings.get.GetSettingsResponse; import org.opensearch.client.Client; import org.opensearch.common.settings.Settings; +import org.opensearch.core.action.ActionListener; public class OpenSearchClient { private final Client client; @@ -20,13 +21,13 @@ public OpenSearchClient(Client client) { this.client = client; } - public Settings getIndexSettings(String indexName, String[] settingNames) { + public void getIndexSettings(String indexName, String[] settingNames, ActionListener settingsListener) { GetSettingsRequest getSettingsRequest = new GetSettingsRequest() .indices(indexName); if (settingNames != null && settingNames.length > 0) { getSettingsRequest.names(settingNames); } - GetSettingsResponse getSettingsResponse = client.execute(GetSettingsAction.INSTANCE, getSettingsRequest).actionGet(); - return getSettingsResponse.getIndexToSettings().get(indexName); + ActionListener responseListener = ActionListener.map(settingsListener, r -> r.getIndexToSettings().get(indexName)); + client.execute(GetSettingsAction.INSTANCE, getSettingsRequest, responseListener); } } diff --git a/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/configuration/ConfigurationUtils.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/configuration/ConfigurationUtils.java new file mode 100644 index 0000000..44081d6 --- /dev/null +++ b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/configuration/ConfigurationUtils.java @@ -0,0 +1,99 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.configuration; + +import org.opensearch.action.search.SearchRequest; +import org.opensearch.common.settings.Settings; +import org.opensearch.search.SearchExtBuilder; +import org.opensearch.search.relevance.transformer.ResultTransformer; + +import java.util.ArrayList; +import java.util.Comparator; +import java.util.List; +import java.util.Map; +import java.util.stream.Collectors; + +import static org.opensearch.search.relevance.configuration.Constants.RESULT_TRANSFORMER_SETTING_PREFIX; + +public class ConfigurationUtils { + + /** + * Get result transformer configurations from Search Request + * + * @param settings all index settings configured for this plugin + * @param resultTransformerMap map of transformed results + * @return ordered and validated list of result transformers, empty list if not specified + */ + public static List getResultTransformersFromIndexConfiguration(Settings settings, + Map resultTransformerMap) { + List indexLevelConfigs = new ArrayList<>(); + + if (settings != null) { + if (settings.getGroups(RESULT_TRANSFORMER_SETTING_PREFIX) != null) { + for (Map.Entry transformerSettings : settings.getGroups(RESULT_TRANSFORMER_SETTING_PREFIX).entrySet()) { + if (resultTransformerMap.containsKey(transformerSettings.getKey())) { + ResultTransformer transformer = resultTransformerMap.get(transformerSettings.getKey()); + indexLevelConfigs.add(transformer.getConfigurationFactory().configure(transformerSettings.getValue())); + } + } + } + } + + return reorderAndValidateConfigs(indexLevelConfigs); + } + + /** + * Get result transformer configurations from Search Request + * + * @param searchRequest input request + * @return ordered and validated list of result transformers, empty list if not specified + */ + public static List getResultTransformersFromRequestConfiguration( + final SearchRequest searchRequest) { + + // Fetch result transformers specified in request + SearchConfigurationExtBuilder requestLevelSearchConfiguration = null; + if (searchRequest.source() != null && searchRequest.source().ext() != null && !searchRequest.source().ext().isEmpty()) { + // Filter ext builders by name + List extBuilders = searchRequest.source().ext().stream() + .filter(searchExtBuilder -> SearchConfigurationExtBuilder.NAME.equals(searchExtBuilder.getWriteableName())) + .collect(Collectors.toList()); + if (!extBuilders.isEmpty()) { + requestLevelSearchConfiguration = (SearchConfigurationExtBuilder) extBuilders.get(0); + } + } + + List requestLevelConfigs = new ArrayList<>(); + if (requestLevelSearchConfiguration != null) { + requestLevelConfigs = reorderAndValidateConfigs(requestLevelSearchConfiguration.getResultTransformers()); + } + return requestLevelConfigs; + } + + /** + * Sort configurations in ascending order of invocation, and validate + * + * @param configs list of result transformer configurations + * @return ordered and validated list of result transformers + */ + public static List reorderAndValidateConfigs( + final List configs) throws IllegalArgumentException { + + // Sort + configs.sort(Comparator.comparingInt(ResultTransformerConfiguration::getOrder)); + + for (int i = 0; i < configs.size(); ++i) { + if (configs.get(i).getOrder() != (i + 1)) { + throw new IllegalArgumentException("Expected order [" + (i + 1) + "] for transformer [" + + configs.get(i).getTransformerName() + "], but found [" + configs.get(i).getOrder() + "]"); + } + } + + return configs; + } +} diff --git a/src/main/java/org/opensearch/search/relevance/configuration/Constants.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/configuration/Constants.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/configuration/Constants.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/configuration/Constants.java diff --git a/src/main/java/org/opensearch/search/relevance/configuration/ResultTransformerConfiguration.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/configuration/ResultTransformerConfiguration.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/configuration/ResultTransformerConfiguration.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/configuration/ResultTransformerConfiguration.java diff --git a/src/main/java/org/opensearch/search/relevance/configuration/ResultTransformerConfigurationFactory.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/configuration/ResultTransformerConfigurationFactory.java similarity index 89% rename from src/main/java/org/opensearch/search/relevance/configuration/ResultTransformerConfigurationFactory.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/configuration/ResultTransformerConfigurationFactory.java index a9f3085..aece3a7 100644 --- a/src/main/java/org/opensearch/search/relevance/configuration/ResultTransformerConfigurationFactory.java +++ b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/configuration/ResultTransformerConfigurationFactory.java @@ -7,7 +7,7 @@ */ package org.opensearch.search.relevance.configuration; -import org.opensearch.common.io.stream.StreamInput; +import org.opensearch.core.common.io.stream.StreamInput; import org.opensearch.common.settings.Settings; import org.opensearch.core.xcontent.XContentParser; @@ -33,7 +33,7 @@ public interface ResultTransformerConfigurationFactory { /** * Build configuration from a serialized stream. - * @param streamInput a {@link org.opensearch.common.io.stream.Writeable} serialized representation of transformer + * @param streamInput a {@link org.opensearch.core.common.io.stream.Writeable} serialized representation of transformer * configuration. * @return configuration the deserialized transformer configuration. */ diff --git a/src/main/java/org/opensearch/search/relevance/configuration/SearchConfigurationExtBuilder.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/configuration/SearchConfigurationExtBuilder.java similarity index 97% rename from src/main/java/org/opensearch/search/relevance/configuration/SearchConfigurationExtBuilder.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/configuration/SearchConfigurationExtBuilder.java index ca082eb..53fea0e 100644 --- a/src/main/java/org/opensearch/search/relevance/configuration/SearchConfigurationExtBuilder.java +++ b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/configuration/SearchConfigurationExtBuilder.java @@ -7,9 +7,9 @@ */ package org.opensearch.search.relevance.configuration; -import org.opensearch.common.ParsingException; -import org.opensearch.common.io.stream.StreamInput; -import org.opensearch.common.io.stream.StreamOutput; +import org.opensearch.core.common.ParsingException; +import org.opensearch.core.common.io.stream.StreamInput; +import org.opensearch.core.common.io.stream.StreamOutput; import org.opensearch.core.ParseField; import org.opensearch.core.xcontent.XContentBuilder; import org.opensearch.core.xcontent.XContentParser; diff --git a/src/main/java/org/opensearch/search/relevance/configuration/TransformerConfiguration.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/configuration/TransformerConfiguration.java similarity index 94% rename from src/main/java/org/opensearch/search/relevance/configuration/TransformerConfiguration.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/configuration/TransformerConfiguration.java index 9178831..5b1099f 100644 --- a/src/main/java/org/opensearch/search/relevance/configuration/TransformerConfiguration.java +++ b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/configuration/TransformerConfiguration.java @@ -10,7 +10,7 @@ import static org.opensearch.search.relevance.configuration.Constants.ORDER; import static org.opensearch.search.relevance.configuration.Constants.PROPERTIES; -import org.opensearch.common.io.stream.Writeable; +import org.opensearch.core.common.io.stream.Writeable; import org.opensearch.core.ParseField; import org.opensearch.core.xcontent.ToXContentObject; diff --git a/src/main/java/org/opensearch/search/relevance/transformer/ResultTransformer.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/ResultTransformer.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/transformer/ResultTransformer.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/ResultTransformer.java diff --git a/src/main/java/org/opensearch/search/relevance/transformer/TransformerType.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/TransformerType.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/transformer/TransformerType.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/TransformerType.java diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/KendraIntelligentRanker.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/KendraIntelligentRanker.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/KendraIntelligentRanker.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/KendraIntelligentRanker.java diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraClientSettings.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraClientSettings.java similarity index 98% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraClientSettings.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraClientSettings.java index 44fb113..523aa99 100644 --- a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraClientSettings.java +++ b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraClientSettings.java @@ -21,7 +21,7 @@ import org.apache.logging.log4j.LogManager; import org.apache.logging.log4j.Logger; -import org.opensearch.common.settings.SecureString; +import org.opensearch.core.common.settings.SecureString; import org.opensearch.common.settings.Settings; import org.opensearch.common.settings.SettingsException; diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraHttpClient.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraHttpClient.java similarity index 96% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraHttpClient.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraHttpClient.java index 64c2a82..11b44dd 100644 --- a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraHttpClient.java +++ b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraHttpClient.java @@ -30,13 +30,12 @@ import java.io.ByteArrayInputStream; import java.io.Closeable; -import java.io.IOException; import java.net.URI; import java.nio.charset.StandardCharsets; import java.security.AccessController; import java.security.PrivilegedAction; -import org.apache.commons.lang3.StringUtils; +import org.opensearch.core.common.Strings; import org.opensearch.search.relevance.transformer.kendraintelligentranking.model.dto.RescoreRequest; import org.opensearch.search.relevance.transformer.kendraintelligentranking.model.dto.RescoreResult; @@ -56,6 +55,7 @@ public class KendraHttpClient implements Closeable { private final ObjectMapper objectMapper = new ObjectMapper() .configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, false); + @SuppressWarnings("removal") public KendraHttpClient(KendraClientSettings clientSettings) { serviceEndpoint = clientSettings.getServiceEndpoint(); executionPlanId = clientSettings.getExecutionPlanId(); @@ -103,6 +103,7 @@ public KendraHttpClient(KendraClientSettings clientSettings) { } } + @SuppressWarnings({ "deprecation", "removal" }) public RescoreResult rescore(RescoreRequest rescoreRequest) { return AccessController.doPrivileged((PrivilegedAction) () -> { try { @@ -132,11 +133,11 @@ public URI buildRescoreURI() { } public boolean isValid() { - return StringUtils.isNotEmpty(serviceEndpoint) && StringUtils.isNotEmpty(executionPlanId); + return !Strings.isNullOrEmpty(serviceEndpoint) && !Strings.isNullOrEmpty(executionPlanId); } @Override - public void close() throws IOException { + public void close() { if (amazonHttpClient != null) { amazonHttpClient.shutdown(); } diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/SimpleAwsErrorHandler.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/SimpleAwsErrorHandler.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/SimpleAwsErrorHandler.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/SimpleAwsErrorHandler.java diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/SimpleResponseHandler.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/SimpleResponseHandler.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/SimpleResponseHandler.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/SimpleResponseHandler.java diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/Constants.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/Constants.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/Constants.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/Constants.java diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankerSettings.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankerSettings.java similarity index 98% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankerSettings.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankerSettings.java index e5f4f23..062a221 100644 --- a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankerSettings.java +++ b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankerSettings.java @@ -12,7 +12,7 @@ import java.util.List; import java.util.function.Function; import org.opensearch.common.settings.SecureSetting; -import org.opensearch.common.settings.SecureString; +import org.opensearch.core.common.settings.SecureString; import org.opensearch.common.settings.Setting; import org.opensearch.common.settings.Setting.Property; diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfiguration.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfiguration.java similarity index 97% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfiguration.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfiguration.java index 3759af0..2842d11 100644 --- a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfiguration.java +++ b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfiguration.java @@ -7,10 +7,10 @@ */ package org.opensearch.search.relevance.transformer.kendraintelligentranking.configuration; -import org.opensearch.common.ParsingException; -import org.opensearch.common.io.stream.StreamInput; -import org.opensearch.common.io.stream.StreamOutput; -import org.opensearch.common.io.stream.Writeable; +import org.opensearch.core.common.ParsingException; +import org.opensearch.core.common.io.stream.StreamInput; +import org.opensearch.core.common.io.stream.StreamOutput; +import org.opensearch.core.common.io.stream.Writeable; import org.opensearch.common.settings.Settings; import org.opensearch.core.ParseField; import org.opensearch.core.xcontent.ObjectParser; diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfigurationFactory.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfigurationFactory.java similarity index 96% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfigurationFactory.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfigurationFactory.java index 7659e9f..080a8c2 100644 --- a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfigurationFactory.java +++ b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfigurationFactory.java @@ -7,7 +7,7 @@ */ package org.opensearch.search.relevance.transformer.kendraintelligentranking.configuration; -import org.opensearch.common.io.stream.StreamInput; +import org.opensearch.core.common.io.stream.StreamInput; import org.opensearch.common.settings.Settings; import org.opensearch.core.xcontent.XContentParser; import org.opensearch.search.relevance.configuration.ResultTransformerConfiguration; diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/KendraIntelligentRankingException.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/KendraIntelligentRankingException.java similarity index 91% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/KendraIntelligentRankingException.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/KendraIntelligentRankingException.java index 41e5494..703d32a 100644 --- a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/KendraIntelligentRankingException.java +++ b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/KendraIntelligentRankingException.java @@ -9,7 +9,7 @@ import java.io.IOException; import org.opensearch.OpenSearchException; -import org.opensearch.common.io.stream.StreamInput; +import org.opensearch.core.common.io.stream.StreamInput; public class KendraIntelligentRankingException extends OpenSearchException { public KendraIntelligentRankingException(StreamInput in) throws IOException { diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/PassageScore.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/PassageScore.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/PassageScore.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/PassageScore.java diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/Document.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/Document.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/Document.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/Document.java diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/RescoreRequest.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/RescoreRequest.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/RescoreRequest.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/RescoreRequest.java diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/RescoreResult.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/RescoreResult.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/RescoreResult.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/RescoreResult.java diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/RescoreResultItem.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/RescoreResultItem.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/RescoreResultItem.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/RescoreResultItem.java diff --git a/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/pipeline/KendraRankingResponseProcessor.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/pipeline/KendraRankingResponseProcessor.java new file mode 100644 index 0000000..26e596c --- /dev/null +++ b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/pipeline/KendraRankingResponseProcessor.java @@ -0,0 +1,179 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ + +package org.opensearch.search.relevance.transformer.kendraintelligentranking.pipeline; + +import org.apache.logging.log4j.LogManager; +import org.apache.logging.log4j.Logger; +import org.opensearch.action.search.SearchRequest; +import org.opensearch.action.search.SearchResponse; +import org.opensearch.action.search.SearchResponseSections; +import org.opensearch.ingest.ConfigurationUtils; +import org.opensearch.search.SearchHits; +import org.opensearch.search.aggregations.InternalAggregations; +import org.opensearch.search.internal.InternalSearchResponse; +import org.opensearch.search.pipeline.AbstractProcessor; +import org.opensearch.search.pipeline.Processor; +import org.opensearch.search.pipeline.SearchResponseProcessor; +import org.opensearch.search.profile.SearchProfileShardResults; +import org.opensearch.search.relevance.transformer.kendraintelligentranking.KendraIntelligentRanker; +import org.opensearch.search.relevance.transformer.kendraintelligentranking.client.KendraClientSettings; +import org.opensearch.search.relevance.transformer.kendraintelligentranking.client.KendraHttpClient; +import org.opensearch.search.relevance.transformer.kendraintelligentranking.configuration.KendraIntelligentRankingConfiguration; + +import static org.opensearch.search.relevance.transformer.kendraintelligentranking.configuration.Constants.KENDRA_DEFAULT_DOC_LIMIT; + +import java.util.Collections; +import java.util.List; +import java.util.Map; +import java.util.concurrent.TimeUnit; + +/** + * This is a {@link SearchResponseProcessor} that applies kendra intelligence ranking + */ +public class KendraRankingResponseProcessor extends AbstractProcessor implements SearchResponseProcessor { + /** + * key to reference this processor type from a search pipeline + */ + public static final String TYPE = "kendra_ranking"; + private final List titleField; + private final List bodyField; + private final int docLimit; + private final String tag; + private final String description; + private final KendraHttpClient kendraClient; + + private static final Logger logger = LogManager.getLogger(KendraRankingResponseProcessor.class); + + /** + * Constructor that apply configuration for kendra re-ranking + * + * @param tag processor tag + * @param description processor description + * @param ignoreFailure processor ignoreFailure config + * @param titleField titleField applied to kendra re-ranking + * @param bodyField bodyField applied to kendra re-ranking + * @param inputDocLimit docLimit applied to kendra re-ranking + * @param kendraClient kendraClient to connect with kendra + */ + public KendraRankingResponseProcessor(String tag, String description, boolean ignoreFailure, List titleField, List bodyField, Integer inputDocLimit, KendraHttpClient kendraClient) { + super(tag, description, ignoreFailure); + this.titleField = titleField; + this.bodyField = bodyField; + this.tag = tag; + this.description = description; + this.kendraClient = kendraClient; + int docLimit; + if (inputDocLimit == null) { + docLimit = KENDRA_DEFAULT_DOC_LIMIT; + } else { + docLimit = inputDocLimit; + } + this.docLimit = docLimit; + } + + /** + * Gets the type of the processor. + */ + @Override + public String getType() { + return TYPE; + } + + /** + * Gets the tag of a processor. + */ + @Override + public String getTag() { + return tag; + } + + /** + * Gets the description of a processor. + */ + @Override + public String getDescription() { + return description; + } + + + /** + * Transform the response hit and apply kendra re-ranking logic + */ + @Override + public SearchResponse processResponse(SearchRequest request, SearchResponse response) throws Exception { + SearchHits hits = response.getHits(); + + if (hits.getHits().length == 0) { + // Avoid call to re-rank empty results + logger.info("TotalHits = 0. Returning search response without transforming."); + return response; + } + + KendraIntelligentRankingConfiguration.KendraIntelligentRankingProperties properties = new KendraIntelligentRankingConfiguration.KendraIntelligentRankingProperties(bodyField, titleField, docLimit); + KendraIntelligentRankingConfiguration configuration = new KendraIntelligentRankingConfiguration(1, properties); + KendraIntelligentRanker ranker = new KendraIntelligentRanker(this.kendraClient); + SearchRequest processedRequest = ranker.preprocessRequest(request, configuration); + + if (ranker.shouldTransform(processedRequest, configuration)) { + long startTime = System.nanoTime(); + SearchHits reRankedSearchHits = ranker.transform(hits, processedRequest, configuration); + long timeTookMillis = TimeUnit.NANOSECONDS.toMillis(System.nanoTime() - startTime); + + final SearchResponseSections internalResponse = new InternalSearchResponse(reRankedSearchHits, + (InternalAggregations) response.getAggregations(), response.getSuggest(), + new SearchProfileShardResults(response.getProfileResults()), response.isTimedOut(), + response.isTerminatedEarly(), response.getNumReducePhases()); + + final SearchResponse newResponse = new SearchResponse(internalResponse, response.getScrollId(), + response.getTotalShards(), response.getSuccessfulShards(), + response.getSkippedShards(), timeTookMillis, response.getShardFailures(), + response.getClusters()); + logger.info("kendra ranking processor took " + timeTookMillis + " ms"); + return newResponse; + } else + return response; + } + + /** + * This is a factor that creates the KendraRankingResponseProcessor + */ + public static final class Factory implements Processor.Factory { + + private final KendraClientSettings clientSettings; + + /** + * Constructor for factory + * @param kendraClientSettings credentials to create kendra client + */ + public Factory(KendraClientSettings kendraClientSettings) { + this.clientSettings = kendraClientSettings; + } + + public KendraRankingResponseProcessor create( + Map> processorFactories, + String tag, + String description, + boolean ignoreFailure, + Map config, + PipelineContext pipelineContext + ) throws Exception { + List titleField = Collections.singletonList(ConfigurationUtils.readOptionalStringProperty(TYPE, tag, config, "title_field")); + List bodyField = Collections.singletonList(ConfigurationUtils.readStringProperty(TYPE, tag, config, "body_field")); + String inputDocLimit = ConfigurationUtils.readOptionalStringOrIntProperty(TYPE, tag, config, "doc_limit"); + KendraHttpClient kendraClient = new KendraHttpClient(this.clientSettings); + int docLimit; + if (inputDocLimit == null) { + docLimit = KENDRA_DEFAULT_DOC_LIMIT; + } else { + docLimit = Integer.parseInt(inputDocLimit); + } + return new KendraRankingResponseProcessor(tag, description, ignoreFailure, titleField, bodyField, docLimit, kendraClient); + } + } +} diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/BM25Scorer.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/BM25Scorer.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/BM25Scorer.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/BM25Scorer.java diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/PassageGenerator.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/PassageGenerator.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/PassageGenerator.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/PassageGenerator.java diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/QueryParser.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/QueryParser.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/QueryParser.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/QueryParser.java diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/SentenceSplitter.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/SentenceSplitter.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/SentenceSplitter.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/SentenceSplitter.java diff --git a/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/TextTokenizer.java b/amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/TextTokenizer.java similarity index 100% rename from src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/TextTokenizer.java rename to amazon-kendra-intelligent-ranking/src/main/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/TextTokenizer.java diff --git a/src/main/plugin-metadata/plugin-security.policy b/amazon-kendra-intelligent-ranking/src/main/plugin-metadata/plugin-security.policy similarity index 100% rename from src/main/plugin-metadata/plugin-security.policy rename to amazon-kendra-intelligent-ranking/src/main/plugin-metadata/plugin-security.policy diff --git a/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/AmazonKendraIntelligentRankingPluginIT.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/AmazonKendraIntelligentRankingPluginIT.java new file mode 100644 index 0000000..cc13e70 --- /dev/null +++ b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/AmazonKendraIntelligentRankingPluginIT.java @@ -0,0 +1,28 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance; + +import org.apache.hc.core5.http.ParseException; +import org.apache.hc.core5.http.io.entity.EntityUtils; +import org.opensearch.client.Request; +import org.opensearch.client.Response; +import org.opensearch.test.rest.OpenSearchRestTestCase; + +import java.io.IOException; + +public class AmazonKendraIntelligentRankingPluginIT extends OpenSearchRestTestCase { + + public void testPluginInstalled() throws IOException, ParseException { + Response response = client().performRequest(new Request("GET", "/_cat/plugins")); + String body = EntityUtils.toString(response.getEntity()); + + logger.info("response body: {}", body); + assertNotNull(body); + assertTrue(body.contains("amazon-kendra-intelligent-ranking")); + } +} diff --git a/src/test/java/org/opensearch/search/relevance/SearchRelevanceTests.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/SearchRelevanceTests.java similarity index 100% rename from src/test/java/org/opensearch/search/relevance/SearchRelevanceTests.java rename to amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/SearchRelevanceTests.java diff --git a/src/test/java/org/opensearch/search/relevance/actionfilter/SearchActionFilterTests.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/actionfilter/SearchActionFilterTests.java similarity index 95% rename from src/test/java/org/opensearch/search/relevance/actionfilter/SearchActionFilterTests.java rename to amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/actionfilter/SearchActionFilterTests.java index 1af2da7..496adbf 100644 --- a/src/test/java/org/opensearch/search/relevance/actionfilter/SearchActionFilterTests.java +++ b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/actionfilter/SearchActionFilterTests.java @@ -9,8 +9,6 @@ import org.apache.lucene.search.TotalHits; import org.mockito.Mockito; -import org.opensearch.action.ActionFuture; -import org.opensearch.action.ActionListener; import org.opensearch.action.admin.indices.settings.get.GetSettingsAction; import org.opensearch.action.admin.indices.settings.get.GetSettingsRequest; import org.opensearch.action.admin.indices.settings.get.GetSettingsResponse; @@ -25,11 +23,11 @@ import org.opensearch.action.search.ShardSearchFailure; import org.opensearch.action.support.ActionFilterChain; import org.opensearch.client.Client; -import org.opensearch.common.bytes.BytesReference; -import org.opensearch.common.collect.ImmutableOpenMap; +import org.opensearch.core.action.ActionListener; +import org.opensearch.core.common.bytes.BytesReference; import org.opensearch.common.document.DocumentField; -import org.opensearch.common.io.stream.StreamInput; -import org.opensearch.common.io.stream.StreamOutput; +import org.opensearch.core.common.io.stream.StreamInput; +import org.opensearch.core.common.io.stream.StreamOutput; import org.opensearch.common.settings.Setting; import org.opensearch.common.settings.Settings; import org.opensearch.common.xcontent.json.JsonXContent; @@ -57,7 +55,8 @@ import static org.mockito.ArgumentMatchers.any; import static org.mockito.ArgumentMatchers.eq; -import static org.mockito.Mockito.when; +import static org.mockito.Mockito.doAnswer; +import static org.mockito.Mockito.mock; public class SearchActionFilterTests extends OpenSearchTestCase { @@ -116,21 +115,19 @@ public void testIgnoresSearchRequestOnMultipleIndices() { private static Client buildMockClient(String indexName, Settings... settings) { Client client = Mockito.mock(Client.class); - ActionFuture mockGetSettingsFuture = Mockito.mock(ActionFuture.class); Settings.Builder settingsBuilder = Settings.builder(); for (Settings settingsEntry : settings) { settingsBuilder.put(settingsEntry); } Settings settingsObj = settingsBuilder.build(); - ImmutableOpenMap indexSettingsMap = ImmutableOpenMap.builder() - .fPut(indexName, settingsObj) - .build(); - ImmutableOpenMap emptyMap = ImmutableOpenMap.builder().build(); - GetSettingsResponse getSettingsResponse = new GetSettingsResponse(indexSettingsMap, emptyMap); - when(mockGetSettingsFuture.actionGet()).thenReturn(getSettingsResponse); - when(client.execute(eq(GetSettingsAction.INSTANCE), any(GetSettingsRequest.class))) - .thenReturn(mockGetSettingsFuture); + Map indexSettingsMap = Map.of(indexName, settingsObj); + GetSettingsResponse getSettingsResponse = new GetSettingsResponse(indexSettingsMap, Collections.emptyMap()); + doAnswer(invocation -> { + ActionListener responseListener = invocation.getArgument(2); + responseListener.onResponse(getSettingsResponse); + return null; + }).when(client).execute(eq(GetSettingsAction.INSTANCE), any(GetSettingsRequest.class), any(ActionListener.class)); return client; } @@ -149,7 +146,8 @@ public void testOperatesOnSingleIndexWithNoTransformers() { AtomicBoolean proceedCalled = new AtomicBoolean(false); ActionFilterChain searchFilterChain = (task1, action, request, listener) -> proceedCalled.set(true); - searchActionFilter.apply(task, SearchAction.NAME, searchRequest, null, searchFilterChain); + ActionListener mockListener = mock(ActionListener.class); + searchActionFilter.apply(task, SearchAction.NAME, searchRequest, mockListener, searchFilterChain); assertTrue(proceedCalled.get()); } @@ -267,7 +265,8 @@ public void testTransformerDoesNotRunWhenNotEnabled() { AtomicBoolean proceedCalled = new AtomicBoolean(false); ActionFilterChain searchFilterChain = (task1, action, request, listener) -> proceedCalled.set(true); - searchActionFilter.apply(task, SearchAction.NAME, searchRequest, null, searchFilterChain); + ActionListener mockListener = mock(ActionListener.class); + searchActionFilter.apply(task, SearchAction.NAME, searchRequest, mockListener, searchFilterChain); assertTrue(proceedCalled.get()); // We should try to check for index-level settings assertTrue(mockTransformer.getTransformerSettingsWasCalled); diff --git a/src/test/java/org/opensearch/search/relevance/configuration/SearchConfigurationExtBuilderTests.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/configuration/SearchConfigurationExtBuilderTests.java similarity index 96% rename from src/test/java/org/opensearch/search/relevance/configuration/SearchConfigurationExtBuilderTests.java rename to amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/configuration/SearchConfigurationExtBuilderTests.java index 3952607..362a5fd 100644 --- a/src/test/java/org/opensearch/search/relevance/configuration/SearchConfigurationExtBuilderTests.java +++ b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/configuration/SearchConfigurationExtBuilderTests.java @@ -7,10 +7,10 @@ */ package org.opensearch.search.relevance.configuration; -import org.opensearch.common.bytes.BytesReference; +import org.opensearch.core.common.bytes.BytesReference; import org.opensearch.common.io.stream.BytesStreamOutput; -import org.opensearch.common.io.stream.StreamInput; -import org.opensearch.common.io.stream.StreamOutput; +import org.opensearch.core.common.io.stream.StreamInput; +import org.opensearch.core.common.io.stream.StreamOutput; import org.opensearch.common.settings.Settings; import org.opensearch.common.xcontent.XContentHelper; import org.opensearch.common.xcontent.XContentType; diff --git a/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/KendraIntelligentRankerTests.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/KendraIntelligentRankerTests.java similarity index 94% rename from src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/KendraIntelligentRankerTests.java rename to amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/KendraIntelligentRankerTests.java index 843693b..1d176d8 100644 --- a/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/KendraIntelligentRankerTests.java +++ b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/KendraIntelligentRankerTests.java @@ -8,9 +8,8 @@ package org.opensearch.search.relevance.transformer.kendraintelligentranking; import org.apache.lucene.search.TotalHits; -import org.mockito.Mockito; import org.opensearch.action.search.SearchRequest; -import org.opensearch.common.bytes.BytesReference; +import org.opensearch.core.common.bytes.BytesReference; import org.opensearch.common.settings.Setting; import org.opensearch.common.settings.Settings; import org.opensearch.common.xcontent.json.JsonXContent; @@ -24,13 +23,13 @@ import org.opensearch.search.relevance.configuration.ResultTransformerConfiguration; import org.opensearch.search.relevance.transformer.kendraintelligentranking.client.KendraClientSettings; import org.opensearch.search.relevance.transformer.kendraintelligentranking.client.KendraHttpClient; +import org.opensearch.search.relevance.transformer.kendraintelligentranking.client.KendraIntelligentClientTests; import org.opensearch.search.relevance.transformer.kendraintelligentranking.configuration.KendraIntelligentRankerSettings; import org.opensearch.search.relevance.transformer.kendraintelligentranking.configuration.KendraIntelligentRankingConfiguration; import org.opensearch.search.relevance.transformer.kendraintelligentranking.configuration.KendraIntelligentRankingConfiguration.KendraIntelligentRankingProperties; import org.opensearch.search.relevance.transformer.kendraintelligentranking.model.dto.RescoreRequest; import org.opensearch.search.relevance.transformer.kendraintelligentranking.model.dto.RescoreResult; import org.opensearch.search.relevance.transformer.kendraintelligentranking.model.dto.RescoreResultItem; -import org.opensearch.test.OpenSearchTestCase; import java.io.IOException; import java.util.Collections; @@ -38,23 +37,9 @@ import java.util.Map; import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicReference; -import java.util.function.Function; import java.util.stream.Collectors; -public class KendraIntelligentRankerTests extends OpenSearchTestCase { - private static KendraHttpClient buildMockHttpClient(Function mockRescoreImpl) { - KendraHttpClient kendraHttpClient = Mockito.mock(KendraHttpClient.class); - Mockito.when(kendraHttpClient.isValid()).thenReturn(true); - Mockito.doAnswer(invocation -> { - RescoreRequest rescoreRequest = invocation.getArgument(0); - return mockRescoreImpl.apply(rescoreRequest); - }).when(kendraHttpClient).rescore(Mockito.any(RescoreRequest.class)); - return kendraHttpClient; - } - - private static KendraHttpClient buildMockHttpClient() { - return buildMockHttpClient(r -> new RescoreResult()); - } +public class KendraIntelligentRankerTests extends KendraIntelligentClientTests { public void testGetSettings() { List> settings = new KendraIntelligentRanker(buildMockHttpClient()).getTransformerSettings(); diff --git a/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraClientSettingsTests.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraClientSettingsTests.java similarity index 100% rename from src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraClientSettingsTests.java rename to amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraClientSettingsTests.java diff --git a/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraHttpClientTests.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraHttpClientTests.java similarity index 100% rename from src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraHttpClientTests.java rename to amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraHttpClientTests.java diff --git a/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraIntelligentClientTests.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraIntelligentClientTests.java new file mode 100644 index 0000000..e459f02 --- /dev/null +++ b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/KendraIntelligentClientTests.java @@ -0,0 +1,32 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.kendraintelligentranking.client; + +import org.mockito.Mockito; +import org.opensearch.search.relevance.transformer.kendraintelligentranking.model.dto.RescoreRequest; +import org.opensearch.search.relevance.transformer.kendraintelligentranking.model.dto.RescoreResult; +import org.opensearch.test.OpenSearchTestCase; + +import java.util.function.Function; + +public class KendraIntelligentClientTests extends OpenSearchTestCase { + protected static KendraHttpClient buildMockHttpClient(Function mockRescoreImpl) { + KendraHttpClient kendraHttpClient = Mockito.mock(KendraHttpClient.class); + Mockito.when(kendraHttpClient.isValid()).thenReturn(true); + Mockito.doAnswer(invocation -> { + RescoreRequest rescoreRequest = invocation.getArgument(0); + return mockRescoreImpl.apply(rescoreRequest); + }).when(kendraHttpClient).rescore(Mockito.any(RescoreRequest.class)); + return kendraHttpClient; + } + + protected static KendraHttpClient buildMockHttpClient() { + return buildMockHttpClient(r -> new RescoreResult()); + } + +} diff --git a/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/SimpleAwsErrorHandlerTests.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/SimpleAwsErrorHandlerTests.java similarity index 100% rename from src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/SimpleAwsErrorHandlerTests.java rename to amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/client/SimpleAwsErrorHandlerTests.java diff --git a/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfigurationTests.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfigurationTests.java similarity index 98% rename from src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfigurationTests.java rename to amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfigurationTests.java index 4b7eeaa..48a307c 100644 --- a/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfigurationTests.java +++ b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/configuration/KendraIntelligentRankingConfigurationTests.java @@ -7,7 +7,7 @@ */ package org.opensearch.search.relevance.transformer.kendraintelligentranking.configuration; -import org.opensearch.common.bytes.BytesReference; +import org.opensearch.core.common.bytes.BytesReference; import org.opensearch.common.io.stream.BytesStreamOutput; import org.opensearch.common.settings.Settings; import org.opensearch.common.xcontent.XContentHelper; diff --git a/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/KendraIntelligentRankingExceptionTests.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/KendraIntelligentRankingExceptionTests.java similarity index 100% rename from src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/KendraIntelligentRankingExceptionTests.java rename to amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/KendraIntelligentRankingExceptionTests.java diff --git a/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/DocumentTests.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/DocumentTests.java similarity index 100% rename from src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/DocumentTests.java rename to amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/model/dto/DocumentTests.java diff --git a/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/pipeline/KendraRankingResponseProcessorTests.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/pipeline/KendraRankingResponseProcessorTests.java new file mode 100644 index 0000000..4a42e3d --- /dev/null +++ b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/pipeline/KendraRankingResponseProcessorTests.java @@ -0,0 +1,137 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.kendraintelligentranking.pipeline; + +import org.apache.lucene.search.TotalHits; +import org.opensearch.OpenSearchParseException; +import org.opensearch.action.search.SearchRequest; +import org.opensearch.action.search.SearchResponse; +import org.opensearch.action.search.SearchResponseSections; +import org.opensearch.core.common.bytes.BytesArray; +import org.opensearch.common.document.DocumentField; +import org.opensearch.common.settings.Settings; +import org.opensearch.env.Environment; +import org.opensearch.env.TestEnvironment; +import org.opensearch.index.query.MatchQueryBuilder; +import org.opensearch.index.query.QueryBuilder; +import org.opensearch.search.SearchHit; +import org.opensearch.search.SearchHits; +import org.opensearch.search.builder.SearchSourceBuilder; +import org.opensearch.search.relevance.transformer.kendraintelligentranking.client.KendraClientSettings; +import org.opensearch.search.relevance.transformer.kendraintelligentranking.client.KendraHttpClient; +import org.opensearch.search.relevance.transformer.kendraintelligentranking.client.KendraIntelligentClientTests; + +import java.util.*; + + +public class KendraRankingResponseProcessorTests extends KendraIntelligentClientTests { + private static final String TYPE = "kendra_ranking"; + private Settings settings = buildEnvSettings(Settings.EMPTY); + private Environment env = TestEnvironment.newEnvironment(settings); + + private KendraClientSettings clientSettings = KendraClientSettings.getClientSettings(env.settings()); + + private SearchRequest createRequest() { + QueryBuilder query = new MatchQueryBuilder("body", "value"); + SearchSourceBuilder source = new SearchSourceBuilder().query(query); + return new SearchRequest().source(source); + } + + private SearchResponse createResponse(int size) { + SearchHit[] hits = new SearchHit[size]; + for (int i = 0; i < size; i++) { + Map searchHitFields = new HashMap<>(); + searchHitFields.put("field", new DocumentField("value" + i, Collections.emptyList())); + searchHitFields.put("body", new DocumentField("body" + i, Collections.emptyList())); + hits[i] = new SearchHit(i, "doc " + i, searchHitFields, Collections.emptyMap()); + hits[i].sourceRef(new BytesArray("{ \"field "+ "\" : \"value"+ i + "\" ,\"body"+ "\" : \"body"+ i + "\" }")); + hits[i].score(i); + } + SearchHits searchHits = new SearchHits(hits, new TotalHits(size * 2L, TotalHits.Relation.EQUAL_TO), size); + SearchResponseSections searchResponseSections = new SearchResponseSections(searchHits, null, null, false, false, null, 0); + return new SearchResponse(searchResponseSections, null, 1, 1, 0, 10, null, null); + } + + public void testFactory() throws Exception { + + KendraRankingResponseProcessor.Factory factory = new KendraRankingResponseProcessor.Factory( this.clientSettings); + + //test create without title field, expect exceptions + expectThrows(OpenSearchParseException.class, () -> factory.create( + Collections.emptyMap(), + null, + null, + false, + Collections.emptyMap(), + null + )); + + //test create with all fields + List titleField= new ArrayList<>(); + titleField.add("field"); + Map configuration = new HashMap<>(); + configuration.put("title_field","field"); + configuration.put("body_field","body"); + configuration.put("doc_limit","500"); + KendraRankingResponseProcessor processorWithAllFields = factory.create(Collections.emptyMap(),"tmp0","testingAllFields", false, configuration,null); + assertEquals(TYPE, processorWithAllFields.getType()); + assertEquals("tmp0", processorWithAllFields.getTag()); + assertEquals("testingAllFields", processorWithAllFields.getDescription()); + + //test create with required field + Map shortConfiguration = new HashMap<>(); + shortConfiguration.put("body_field","body"); + KendraRankingResponseProcessor processorWithOneFields = factory.create(Collections.emptyMap(),"tmp1","testingBodyField", false, shortConfiguration, null); + assertEquals(TYPE, processorWithOneFields.getType()); + assertEquals("tmp1", processorWithOneFields.getTag()); + assertEquals("testingBodyField", processorWithOneFields.getDescription()); + + //test create with null doc_limit field + Map nullDocLimitConfiguration = new HashMap<>(); + nullDocLimitConfiguration.put("body_field","body"); + nullDocLimitConfiguration.put("doc_limit",null); + KendraRankingResponseProcessor processorWithNullDocLimit = factory.create(Collections.emptyMap(),"tmp2","testingNullDocLimit", false, nullDocLimitConfiguration, null ); + assertEquals(TYPE, processorWithNullDocLimit.getType()); + assertEquals("tmp2", processorWithNullDocLimit.getTag()); + assertEquals("testingNullDocLimit", processorWithNullDocLimit.getDescription()); + + //test create with null title field + Map nullTitleConfiguration = new HashMap<>(); + nullTitleConfiguration.put("body_field","body"); + nullTitleConfiguration.put("title_field",null); + KendraRankingResponseProcessor processorWithNullTitleField = factory.create(Collections.emptyMap(),"tmp3","testingNullTitleField", false, nullTitleConfiguration, null); + assertEquals(TYPE, processorWithNullTitleField.getType()); + assertEquals("tmp3", processorWithNullTitleField.getTag()); + assertEquals("testingNullTitleField", processorWithNullTitleField.getDescription()); + + } + public void testRankingResponse() throws Exception { + KendraHttpClient kendraClient = buildMockHttpClient(); + List titleField = new ArrayList<>(); + titleField.add("field"); + List bodyField = new ArrayList<>(); + bodyField.add("body"); + + //test response with titleField, bodyField and docLimit + KendraRankingResponseProcessor processorWtOptionalConfig = new KendraRankingResponseProcessor(null,null,false, titleField,bodyField,500,kendraClient); + int size = 5; + SearchResponse reRankedResponse0 = processorWtOptionalConfig.processResponse(createRequest(),createResponse(size)); + assertEquals(size,reRankedResponse0.getHits().getHits().length); + + //test response with null doc limit + KendraRankingResponseProcessor processorWtTwoConfig = new KendraRankingResponseProcessor(null,null,false, titleField,bodyField,null,kendraClient); + SearchResponse reRankedResponse1 = processorWtTwoConfig.processResponse(createRequest(),createResponse(size)); + assertEquals(size,reRankedResponse1.getHits().getHits().length); + + //test response with null doc limit and null title field + KendraRankingResponseProcessor processorWtOneConfig = new KendraRankingResponseProcessor(null,null,false,null,bodyField,null,kendraClient); + SearchResponse reRankedResponse2 = processorWtOneConfig.processResponse(createRequest(),createResponse(size)); + assertEquals(size,reRankedResponse2.getHits().getHits().length); + + } +} diff --git a/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/BM25ScorerTests.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/BM25ScorerTests.java similarity index 100% rename from src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/BM25ScorerTests.java rename to amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/BM25ScorerTests.java diff --git a/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/PassageGeneratorTests.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/PassageGeneratorTests.java similarity index 100% rename from src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/PassageGeneratorTests.java rename to amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/PassageGeneratorTests.java diff --git a/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/QueryParserTests.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/QueryParserTests.java similarity index 100% rename from src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/QueryParserTests.java rename to amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/QueryParserTests.java diff --git a/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/SentenceSplitterTests.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/SentenceSplitterTests.java similarity index 100% rename from src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/SentenceSplitterTests.java rename to amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/SentenceSplitterTests.java diff --git a/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/TextTokenizerTests.java b/amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/TextTokenizerTests.java similarity index 100% rename from src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/TextTokenizerTests.java rename to amazon-kendra-intelligent-ranking/src/test/java/org/opensearch/search/relevance/transformer/kendraintelligentranking/preprocess/TextTokenizerTests.java diff --git a/amazon-kendra-intelligent-ranking/src/yamlRestTest/java/org/opensearch/search/relevance/AmazonKendraIntelligentRankingClientYamlTestSuiteIT.java b/amazon-kendra-intelligent-ranking/src/yamlRestTest/java/org/opensearch/search/relevance/AmazonKendraIntelligentRankingClientYamlTestSuiteIT.java new file mode 100644 index 0000000..ae37662 --- /dev/null +++ b/amazon-kendra-intelligent-ranking/src/yamlRestTest/java/org/opensearch/search/relevance/AmazonKendraIntelligentRankingClientYamlTestSuiteIT.java @@ -0,0 +1,26 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance; + +import com.carrotsearch.randomizedtesting.annotations.Name; +import com.carrotsearch.randomizedtesting.annotations.ParametersFactory; +import org.opensearch.test.rest.yaml.ClientYamlTestCandidate; +import org.opensearch.test.rest.yaml.OpenSearchClientYamlSuiteTestCase; + + +public class AmazonKendraIntelligentRankingClientYamlTestSuiteIT extends OpenSearchClientYamlSuiteTestCase { + + public AmazonKendraIntelligentRankingClientYamlTestSuiteIT(@Name("yaml") ClientYamlTestCandidate testCandidate) { + super(testCandidate); + } + + @ParametersFactory + public static Iterable parameters() throws Exception { + return OpenSearchClientYamlSuiteTestCase.createParameters(); + } +} diff --git a/amazon-kendra-intelligent-ranking/src/yamlRestTest/resources/rest-api-spec/test/10_basic.yml b/amazon-kendra-intelligent-ranking/src/yamlRestTest/resources/rest-api-spec/test/10_basic.yml new file mode 100644 index 0000000..af5e5fe --- /dev/null +++ b/amazon-kendra-intelligent-ranking/src/yamlRestTest/resources/rest-api-spec/test/10_basic.yml @@ -0,0 +1,17 @@ +"Test that the plugin is loaded in OpenSearch": + - do: + cat.plugins: + local: true + h: component + + - match: + $body: /^opensearch-amazon-kendra-intelligent-ranking-\d+.\d+.\d+.\d+\n$/ + + - do: + indices.create: + index: test + + - do: + search: + index: test + body: { } diff --git a/amazon-personalize-ranking/build.gradle b/amazon-personalize-ranking/build.gradle new file mode 100644 index 0000000..527d985 --- /dev/null +++ b/amazon-personalize-ranking/build.gradle @@ -0,0 +1,129 @@ +import org.opensearch.gradle.test.RestIntegTestTask + +apply plugin: 'java' +apply plugin: 'idea' +apply plugin: 'opensearch.opensearchplugin' +apply plugin: 'opensearch.yaml-rest-test' +apply plugin: 'jacoco' + +group = 'org.opensearch' + +def pluginName = 'amazon-personalize-ranking' +def pluginDescription = 'Rerank search results using Amazon Personalize' +def projectPath = 'org.opensearch' +def pathToPlugin = 'search.relevance' +def pluginClassName = 'AmazonPersonalizeRankingPlugin' + + +opensearchplugin { + name "opensearch-${pluginName}-${plugin_version}.0" + version "${plugin_version}" + description pluginDescription + classname "${projectPath}.${pathToPlugin}.${pluginClassName}" + licenseFile rootProject.file('LICENSE') + noticeFile rootProject.file('NOTICE') +} + +java { + targetCompatibility = JavaVersion.VERSION_21 + sourceCompatibility = JavaVersion.VERSION_21 +} + +// This requires an additional Jar not published as part of build-tools +loggerUsageCheck.enabled = false + +// No need to validate pom, as we do not upload to maven/sonatype +validateNebulaPom.enabled = false + +buildscript { + repositories { + mavenLocal() + maven { url "https://aws.oss.sonatype.org/content/repositories/snapshots" } + mavenCentral() + maven { url "https://plugins.gradle.org/m2/" } + } + + dependencies { + classpath "org.opensearch.gradle:build-tools:${opensearch_version}" + } +} + +repositories { + mavenLocal() + maven { url "https://aws.oss.sonatype.org/content/repositories/snapshots" } + mavenCentral() + maven { url "https://plugins.gradle.org/m2/" } +} + +dependencies { + implementation 'org.apache.httpcomponents:httpclient:4.5.14' + implementation 'org.apache.httpcomponents:httpcore:4.4.16' + implementation 'com.fasterxml.jackson.core:jackson-databind:2.18.2' + implementation 'com.fasterxml.jackson.core:jackson-core:2.18.2' + implementation 'com.fasterxml.jackson.core:jackson-annotations:2.18.2' + implementation 'com.amazonaws:aws-java-sdk-sts:1.12.300' + implementation 'com.amazonaws:aws-java-sdk-core:1.12.300' + implementation 'com.amazonaws:aws-java-sdk-personalizeruntime:1.12.300' + implementation 'commons-logging:commons-logging:1.2' +} + + +allprojects { + plugins.withId('jacoco') { + jacoco.toolVersion = '0.8.9' + } +} + + +test { + include '**/*Tests.class' + finalizedBy jacocoTestReport +} + +task integTest(type: RestIntegTestTask) { + description = "Run tests against a cluster" + testClassesDirs = sourceSets.test.output.classesDirs + classpath = sourceSets.test.runtimeClasspath +} +tasks.named("check").configure { dependsOn(integTest) } + +integTest { + // The --debug-jvm command-line option makes the cluster debuggable; this makes the tests debuggable + if (System.getProperty("test.debug") != null) { + jvmArgs '-agentlib:jdwp=transport=dt_socket,server=y,suspend=y,address=*:5005' + } +} + +testClusters.integTest { + testDistribution = "ARCHIVE" + + // This installs our plugin into the testClusters + plugin(project.tasks.bundlePlugin.archiveFile) +} + +run { + useCluster testClusters.integTest +} + +sourceSets { + main { + resources { + srcDirs = ["config"] + includes = ["**/*.yml"] + } + } +} + + +jacocoTestReport { + dependsOn test + reports { + xml.required = true + html.required = true + } +} + +// TODO: Enable these checks +dependencyLicenses.enabled = false +thirdPartyAudit.enabled = false +loggerUsageCheck.enabled = false diff --git a/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/AmazonPersonalizeRankingPlugin.java b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/AmazonPersonalizeRankingPlugin.java new file mode 100644 index 0000000..343fae5 --- /dev/null +++ b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/AmazonPersonalizeRankingPlugin.java @@ -0,0 +1,80 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance; + +import org.opensearch.client.Client; +import org.opensearch.cluster.metadata.IndexNameExpressionResolver; +import org.opensearch.cluster.service.ClusterService; +import org.opensearch.common.settings.Setting; +import org.opensearch.core.common.io.stream.NamedWriteableRegistry; +import org.opensearch.core.xcontent.NamedXContentRegistry; +import org.opensearch.env.Environment; +import org.opensearch.env.NodeEnvironment; +import org.opensearch.plugins.Plugin; +import org.opensearch.plugins.SearchPipelinePlugin; +import org.opensearch.plugins.SearchPlugin; +import org.opensearch.repositories.RepositoriesService; +import org.opensearch.script.ScriptService; +import org.opensearch.search.pipeline.Processor; +import org.opensearch.search.pipeline.SearchResponseProcessor; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.PersonalizeRankingResponseProcessor; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.client.PersonalizeClientSettings; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.requestparameter.PersonalizeRequestParametersExtBuilder; +import org.opensearch.threadpool.ThreadPool; +import org.opensearch.watcher.ResourceWatcherService; + +import java.util.ArrayList; +import java.util.Collection; +import java.util.Collections; +import java.util.List; +import java.util.Map; +import java.util.function.Supplier; +import java.util.stream.Collectors; + +public class AmazonPersonalizeRankingPlugin extends Plugin implements SearchPlugin, SearchPipelinePlugin { + + private PersonalizeClientSettings personalizeClientSettings; + + @Override + public List> getSettings() { + // Add settings for other transformers here + return new ArrayList<>(PersonalizeClientSettings.getAllSettings()); + } + + @Override + public Collection createComponents( + Client client, + ClusterService clusterService, + ThreadPool threadPool, + ResourceWatcherService resourceWatcherService, + ScriptService scriptService, + NamedXContentRegistry xContentRegistry, + Environment environment, + NodeEnvironment nodeEnvironment, + NamedWriteableRegistry namedWriteableRegistry, + IndexNameExpressionResolver indexNameExpressionResolver, + Supplier repositoriesServiceSupplier + ) { + this.personalizeClientSettings = PersonalizeClientSettings.getClientSettings(environment.settings()); + + return Collections.emptyList(); + } + + @Override + public List> getSearchExts() { + return List.of( + new SearchPlugin.SearchExtSpec<>(PersonalizeRequestParametersExtBuilder.NAME, + PersonalizeRequestParametersExtBuilder::new, + PersonalizeRequestParametersExtBuilder::parse)); + } + + @Override + public Map> getResponseProcessors(Parameters parameters) { + return Map.of(PersonalizeRankingResponseProcessor.TYPE, new PersonalizeRankingResponseProcessor.Factory(this.personalizeClientSettings)); + } +} \ No newline at end of file diff --git a/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/PersonalizeRankingResponseProcessor.java b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/PersonalizeRankingResponseProcessor.java new file mode 100644 index 0000000..670915c --- /dev/null +++ b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/PersonalizeRankingResponseProcessor.java @@ -0,0 +1,189 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking; + +import com.amazonaws.auth.AWSCredentialsProvider; +import org.apache.logging.log4j.LogManager; +import org.apache.logging.log4j.Logger; +import org.opensearch.action.search.SearchRequest; +import org.opensearch.action.search.SearchResponse; +import org.opensearch.action.search.SearchResponseSections; +import org.opensearch.ingest.ConfigurationUtils; +import org.opensearch.search.SearchHits; +import org.opensearch.search.aggregations.InternalAggregations; +import org.opensearch.search.internal.InternalSearchResponse; +import org.opensearch.search.pipeline.AbstractProcessor; +import org.opensearch.search.pipeline.Processor; +import org.opensearch.search.pipeline.SearchResponseProcessor; +import org.opensearch.search.profile.SearchProfileShardResults; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.client.PersonalizeClient; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.client.PersonalizeClientSettings; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.client.PersonalizeCredentialsProviderFactory; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration.PersonalizeIntelligentRankerConfiguration; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.requestparameter.PersonalizeRequestParameterUtil; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.requestparameter.PersonalizeRequestParameters; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.reranker.PersonalizedRanker; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.reranker.PersonalizedRankerFactory; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.utils.ValidationUtil; + +import java.util.Map; +import java.util.concurrent.TimeUnit; +import java.util.function.BiFunction; + +/** + * This is a {@link SearchResponseProcessor} that applies Personalized intelligent ranking + */ +public class PersonalizeRankingResponseProcessor extends AbstractProcessor implements SearchResponseProcessor { + + private static final Logger logger = LogManager.getLogger(PersonalizeRankingResponseProcessor.class); + + public static final String TYPE = "personalized_search_ranking"; + private final String tag; + private final String description; + private final PersonalizeClient personalizeClient; + private final PersonalizeIntelligentRankerConfiguration rankerConfig; + + /** + * Constructor for Personalize ranking response processor + * + * @param tag processor tag + * @param description processor description + * @param ignoreFailure processor ignoreFailure config + * @param rankerConfig personalize ranker config + * @param client personalize client + */ + public PersonalizeRankingResponseProcessor(String tag, + String description, + boolean ignoreFailure, + PersonalizeIntelligentRankerConfiguration rankerConfig, + PersonalizeClient client) { + super(tag, description, ignoreFailure); + this.tag = tag; + this.description = description; + this.rankerConfig = rankerConfig; + this.personalizeClient = client; + } + + /** + * Transform the response hits by re ranking results using Personalize + * + * @param request Search request + * @param response Search response that needs to be transformed + * @return Transformed search response using personalized re ranking + * @throws Exception Throws exception for any error while processing response + */ + @Override + public SearchResponse processResponse(SearchRequest request, SearchResponse response) throws Exception { + SearchHits hits = response.getHits(); + + if (hits.getHits().length == 0) { + logger.info("TotalHits = 0. Returning search response without applying Personalize transform"); + return response; + } + logger.info("Personalizing search results."); + PersonalizeRequestParameters personalizeRequestParameters = + PersonalizeRequestParameterUtil.getPersonalizeRequestParameters(request); + PersonalizedRankerFactory rankerFactory = new PersonalizedRankerFactory(); + PersonalizedRanker ranker = rankerFactory.getPersonalizedRanker(rankerConfig, personalizeClient); + long startTime = System.nanoTime(); + SearchHits personalizedHits = ranker.rerank(hits, personalizeRequestParameters); + long personalizeTimeTookMs = TimeUnit.NANOSECONDS.toMillis(System.nanoTime() - startTime); + + final SearchResponseSections transformedSearchResponseSections = new InternalSearchResponse(personalizedHits, + (InternalAggregations) response.getAggregations(), response.getSuggest(), + new SearchProfileShardResults(response.getProfileResults()), response.isTimedOut(), + response.isTerminatedEarly(), response.getNumReducePhases()); + + final SearchResponse transformedResponse = new SearchResponse(transformedSearchResponseSections, response.getScrollId(), + response.getTotalShards(), response.getSuccessfulShards(), + response.getSkippedShards(), response.getTook().getMillis() + personalizeTimeTookMs, response.getShardFailures(), + response.getClusters()); + + logger.info("Personalize ranking processor took " + personalizeTimeTookMs + " ms"); + + return transformedResponse; + } + + /** + * Get the type of the processor. + */ + @Override + public String getType() { + return TYPE; + } + + /** + * Get the tag of a processor. + */ + @Override + public String getTag() { + return tag; + } + + /** + * Gets the description of a processor. + */ + @Override + public String getDescription() { + return description; + } + + public static final class Factory implements Processor.Factory { + + private static final String CAMPAIGN_ARN_CONFIG_NAME = "campaign_arn"; + private static final String ITEM_ID_FIELD_CONFIG_NAME = "item_id_field"; + private static final String IAM_ROLE_ARN_CONFIG_NAME = "iam_role_arn"; + private static final String RECIPE_CONFIG_NAME = "recipe"; + private static final String REGION_CONFIG_NAME = "aws_region"; + private static final String WEIGHT_CONFIG_NAME = "weight"; + PersonalizeClientSettings personalizeClientSettings; + private final BiFunction clientBuilder; + + Factory(PersonalizeClientSettings settings, BiFunction clientBuilder) { + this.personalizeClientSettings = settings; + this.clientBuilder = clientBuilder; + } + + public Factory(PersonalizeClientSettings settings) { + this(settings, PersonalizeClient::new); + } + + @Override + public PersonalizeRankingResponseProcessor create(Map> processorFactories, String tag, String description, boolean ignoreFailure, Map config, PipelineContext pipelineContext) { + String personalizeCampaign = ConfigurationUtils.readStringProperty(TYPE, tag, config, CAMPAIGN_ARN_CONFIG_NAME); + String iamRoleArn = ConfigurationUtils.readOptionalStringProperty(TYPE, tag, config, IAM_ROLE_ARN_CONFIG_NAME); + String recipe = ConfigurationUtils.readStringProperty(TYPE, tag, config, RECIPE_CONFIG_NAME); + String itemIdField = ConfigurationUtils.readOptionalStringProperty(TYPE, tag, config, ITEM_ID_FIELD_CONFIG_NAME); + String awsRegion = ConfigurationUtils.readStringProperty(TYPE, tag, config, REGION_CONFIG_NAME); + double weight = ConfigurationUtils.readDoubleProperty(TYPE, tag, config, WEIGHT_CONFIG_NAME); + + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, recipe, itemIdField, awsRegion, weight); + ValidationUtil.validatePersonalizeIntelligentRankerConfiguration(rankerConfig, TYPE, tag); + + final PersonalizeClient personalizeClient; + switch (pipelineContext.getPipelineSource()) { + case SEARCH_REQUEST: + throw new IllegalStateException(TYPE + " processor may not be instantiated as part of a search request. Create a named search pipeline instead."); + case UPDATE_PIPELINE: + AWSCredentialsProvider credentialsProvider = PersonalizeCredentialsProviderFactory.getCredentialsProvider(personalizeClientSettings, iamRoleArn, awsRegion); + personalizeClient = clientBuilder.apply(credentialsProvider, awsRegion); + break; + case VALIDATE_PIPELINE: + default: + personalizeClient = null; // Do not instantiate client on validation + } + return new PersonalizeRankingResponseProcessor(tag, description, ignoreFailure, rankerConfig, personalizeClient); + } + } + + PersonalizeClient getPersonalizeClient() { + // Visible for testing + return personalizeClient; + } +} diff --git a/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeClient.java b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeClient.java new file mode 100644 index 0000000..738acf2 --- /dev/null +++ b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeClient.java @@ -0,0 +1,77 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.client; + +import com.amazonaws.AmazonServiceException; +import com.amazonaws.ClientConfiguration; +import com.amazonaws.auth.AWSCredentialsProvider; +import com.amazonaws.services.personalizeruntime.AmazonPersonalizeRuntime; +import com.amazonaws.services.personalizeruntime.AmazonPersonalizeRuntimeClientBuilder; +import com.amazonaws.services.personalizeruntime.model.GetPersonalizedRankingRequest; +import com.amazonaws.services.personalizeruntime.model.GetPersonalizedRankingResult; + +import java.io.Closeable; +import java.io.IOException; +import java.security.AccessController; +import java.security.PrivilegedAction; + +/** + * Amazon Personalize client implementation for getting personalized ranking + */ +public class PersonalizeClient implements Closeable { + private final AmazonPersonalizeRuntime personalizeRuntime; + private static final String USER_AGENT_PREFIX = "PersonalizeOpenSearchPlugin"; + + /** + * Constructor for Amazon Personalize client + * @param credentialsProvider Credentials to be used for accessing Amazon Personalize + * @param awsRegion AWS region where Amazon Personalize campaign is hosted + */ + public PersonalizeClient(AWSCredentialsProvider credentialsProvider, String awsRegion) { + ClientConfiguration clientConfiguration = AccessController.doPrivileged( + (PrivilegedAction) () -> new ClientConfiguration() + .withUserAgentPrefix(USER_AGENT_PREFIX)); + personalizeRuntime = AccessController.doPrivileged( + (PrivilegedAction) () -> AmazonPersonalizeRuntimeClientBuilder.standard() + .withCredentials(credentialsProvider) + .withRegion(awsRegion) + .withClientConfiguration(clientConfiguration) + .build()); + } + + /** + * Get Personalize runtime client + * @return Personalize runtime client + */ + public AmazonPersonalizeRuntime getPersonalizeRuntime() { + return personalizeRuntime; + } + + /** + * Get Personalized ranking using Personalized runtime client + * @param request Get personalized ranking request + * @return Personalized ranking results + */ + public GetPersonalizedRankingResult getPersonalizedRanking(GetPersonalizedRankingRequest request) { + GetPersonalizedRankingResult result; + try { + result = AccessController.doPrivileged( + (PrivilegedAction) () -> personalizeRuntime.getPersonalizedRanking(request)); + } catch (AmazonServiceException ex) { + throw ex; + } + return result; + } + + @Override + public void close() throws IOException { + if (personalizeRuntime != null) { + personalizeRuntime.shutdown(); + } + } +} \ No newline at end of file diff --git a/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeClientSettings.java b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeClientSettings.java new file mode 100644 index 0000000..1a019b5 --- /dev/null +++ b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeClientSettings.java @@ -0,0 +1,106 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.client; + +import com.amazonaws.auth.AWSCredentials; +import com.amazonaws.auth.BasicAWSCredentials; +import com.amazonaws.auth.BasicSessionCredentials; +import org.apache.logging.log4j.LogManager; +import org.apache.logging.log4j.Logger; +import org.opensearch.core.common.settings.SecureString; +import org.opensearch.common.settings.SecureSetting; +import org.opensearch.common.settings.Setting; +import org.opensearch.common.settings.Settings; +import org.opensearch.common.settings.SettingsException; + +import java.util.Arrays; +import java.util.Collection; + +/** + * Container for personalize client settings such as AWS credentials + */ +public final class PersonalizeClientSettings { + + private static final Logger logger = LogManager.getLogger(PersonalizeClientSettings.class); + + /** + * The access key (ie login id) for connecting to Personalize. + */ + public static final Setting ACCESS_KEY_SETTING = SecureSetting.secureString("personalized_search_ranking.aws.access_key", null); + + /** + * The secret key (ie password) for connecting to Personalize. + */ + public static final Setting SECRET_KEY_SETTING = SecureSetting.secureString("personalized_search_ranking.aws.secret_key", null); + + /** + * The session token for connecting to Personalize. + */ + public static final Setting SESSION_TOKEN_SETTING = SecureSetting.secureString("personalized_search_ranking.aws.session_token", null); + + private final AWSCredentials credentials; + + protected PersonalizeClientSettings(AWSCredentials credentials) { + this.credentials = credentials; + } + + public static Collection> getAllSettings() { + return Arrays.asList( + ACCESS_KEY_SETTING, + SECRET_KEY_SETTING, + SESSION_TOKEN_SETTING + ); + } + + public AWSCredentials getCredentials() { + return credentials; + } + + /** + * Load AWS credentials from open search keystore if available + * @param settings Open search settings + * @return AWS credentials + */ + static AWSCredentials loadCredentials(Settings settings) { + try (SecureString key = ACCESS_KEY_SETTING.get(settings); + SecureString secret = SECRET_KEY_SETTING.get(settings); + SecureString sessionToken = SESSION_TOKEN_SETTING.get(settings)) { + if (key.length() == 0 && secret.length() == 0) { + if (sessionToken.length() > 0) { + throw new SettingsException("Setting [{}] is set but [{}] and [{}] are not", + SESSION_TOKEN_SETTING.getKey(), ACCESS_KEY_SETTING.getKey(), SECRET_KEY_SETTING.getKey()); + } + logger.info("Using either environment variables, system properties or instance profile credentials"); + return null; + } else if (key.length() == 0 || secret.length() == 0) { + throw new SettingsException("One of settings [{}] and [{}] is not set.", + ACCESS_KEY_SETTING.getKey(), SECRET_KEY_SETTING.getKey()); + } else { + final AWSCredentials credentials; + if (sessionToken.length() == 0) { + logger.info("Using basic key/secret credentials"); + credentials = new BasicAWSCredentials(key.toString(), secret.toString()); + } else { + logger.info("Using basic session credentials"); + credentials = new BasicSessionCredentials(key.toString(), secret.toString(), sessionToken.toString()); + } + return credentials; + } + } + } + + /** + * Get Personalize client settings + * @param settings Open search settings + * @return Personalize client settings instance with AWS credentials + */ + public static PersonalizeClientSettings getClientSettings(Settings settings) { + final AWSCredentials credentials = loadCredentials(settings); + return new PersonalizeClientSettings(credentials); + } +} diff --git a/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeCredentialsProviderFactory.java b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeCredentialsProviderFactory.java new file mode 100644 index 0000000..e0e2895 --- /dev/null +++ b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeCredentialsProviderFactory.java @@ -0,0 +1,89 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.client; + +import com.amazonaws.auth.AWSCredentials; +import com.amazonaws.auth.AWSCredentialsProvider; +import com.amazonaws.auth.DefaultAWSCredentialsProviderChain; +import com.amazonaws.auth.AWSStaticCredentialsProvider; +import com.amazonaws.auth.STSAssumeRoleSessionCredentialsProvider; +import com.amazonaws.services.securitytoken.AWSSecurityTokenService; +import com.amazonaws.services.securitytoken.AWSSecurityTokenServiceClientBuilder; +import org.apache.logging.log4j.LogManager; +import org.apache.logging.log4j.Logger; + +import java.security.AccessController; +import java.security.PrivilegedAction; + +/** + * Factory implementation for getting Personalize credentials + */ +public final class PersonalizeCredentialsProviderFactory { + private static final Logger logger = LogManager.getLogger(PersonalizeCredentialsProviderFactory.class); + private static final String ASSUME_ROLE_SESSION_NAME = "OpenSearchPersonalizeIntelligentRankingPluginSession"; + + private PersonalizeCredentialsProviderFactory() { + } + + /** + * Get AWS credentials provider either from static credentials from open search keystore or + * using DefaultAWSCredentialsProviderChain. + * @param clientSettings Personalize client settings + * @return AWS credentials provider for accessing Personalize + */ + static AWSCredentialsProvider getCredentialsProvider(PersonalizeClientSettings clientSettings) { + final AWSCredentialsProvider credentialsProvider; + final AWSCredentials credentials = clientSettings.getCredentials(); + if (credentials == null) { + logger.info("Credentials not present in open search keystore. Using DefaultAWSCredentialsProviderChain for credentials."); + credentialsProvider = AccessController.doPrivileged( + (PrivilegedAction) () -> DefaultAWSCredentialsProviderChain.getInstance()); + } else { + logger.info("Using credentials provided in open search keystore"); + credentialsProvider = AccessController.doPrivileged( + (PrivilegedAction) () -> new AWSStaticCredentialsProvider(credentials)); + } + return credentialsProvider; + } + + /** + * Get AWS credentials provider by assuming IAM role if provided or else + * use static credentials or DefaultAWSCredentialsProviderChain. + * @param clientSettings Personalize client settings + * @param personalizeIAMRole IAM role configuration for accessing Personalize + * @param awsRegion AWS region + * @return AWS credentials provider for accessing Amazon Personalize + */ + public static AWSCredentialsProvider getCredentialsProvider(PersonalizeClientSettings clientSettings, + String personalizeIAMRole, + String awsRegion) { + + final AWSCredentialsProvider credentialsProvider; + AWSCredentialsProvider baseCredentialsProvider = getCredentialsProvider(clientSettings); + + if (personalizeIAMRole != null && !personalizeIAMRole.isBlank()) { + logger.info("Using IAM Role provided to access Personalize."); + // If IAM role ARN was provided in config, then use auto-refreshed role credentials. + credentialsProvider = AccessController.doPrivileged( + (PrivilegedAction) () -> { + AWSSecurityTokenService awsSecurityTokenService = AWSSecurityTokenServiceClientBuilder.standard() + .withCredentials(baseCredentialsProvider) + .withRegion(awsRegion) + .build(); + + return new STSAssumeRoleSessionCredentialsProvider.Builder(personalizeIAMRole, ASSUME_ROLE_SESSION_NAME) + .withStsClient(awsSecurityTokenService) + .build(); + }); + } else { + logger.info("IAM Role for accessing Personalize is not provided."); + credentialsProvider = baseCredentialsProvider; + } + return credentialsProvider; + } +} diff --git a/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/configuration/Constants.java b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/configuration/Constants.java new file mode 100644 index 0000000..40105f9 --- /dev/null +++ b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/configuration/Constants.java @@ -0,0 +1,17 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ + +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration; + +/** + * Constants for Amazon Perosnalize response processor + */ +public class Constants { + public static final String AMAZON_PERSONALIZED_RANKING_RECIPE_NAME = "aws-personalized-ranking"; + public static final String AMAZON_PERSONALIZED_RANKING_V2_RECIPE_NAME = "aws-personalized-ranking-v2"; +} diff --git a/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/configuration/PersonalizeIntelligentRankerConfiguration.java b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/configuration/PersonalizeIntelligentRankerConfiguration.java new file mode 100644 index 0000000..35685de --- /dev/null +++ b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/configuration/PersonalizeIntelligentRankerConfiguration.java @@ -0,0 +1,91 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration; + +/** + * A container for holding Personalize ranker configuration + */ +public class PersonalizeIntelligentRankerConfiguration { + private final String personalizeCampaign; + private final String iamRoleArn; + private final String recipe; + private final String itemIdField; + private final String region; + private final double weight; + + /** + * + * @param personalizeCampaign Personalize campaign + * @param iamRoleArn IAM Role ARN for accessing Personalize campaign + * @param recipe Personalize recipe associated with campaign + * @param itemIdField Item ID field to pick up item id for Personalize input + * @param region AWS region + * @param weight Configurable coefficient to control Personalization of search results + */ + public PersonalizeIntelligentRankerConfiguration(String personalizeCampaign, + String iamRoleArn, + String recipe, + String itemIdField, + String region, + double weight) { + this.personalizeCampaign = personalizeCampaign; + this.iamRoleArn = iamRoleArn; + this.recipe = recipe; + this.itemIdField = itemIdField; + this.region = region; + this.weight = weight; + } + + /** + * Get PErsonalize campaign + * @return Personalize campaign + */ + public String getPersonalizeCampaign() { + return personalizeCampaign; + } + + /** + * Get recipe + * @return Recipe associated with Personalize campaign + */ + public String getRecipe() { + return recipe; + } + + /** + * Get Item ID field + * @return Item ID field + */ + public String getItemIdField() { + return itemIdField; + } + + /** + * Get AWS region + * @return AWS region + */ + public String getRegion() { + return region; + } + + /** + * + * @return weight value + */ + public double getWeight() { + return weight; + } + + /** + * Get IAM role ARN for Personalize campaign + * @return IAM role for accessing Personalize campaign + */ + public String getIamRoleArn() { + return iamRoleArn; + } +} diff --git a/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/requestparameter/PersonalizeRequestParameterUtil.java b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/requestparameter/PersonalizeRequestParameterUtil.java new file mode 100644 index 0000000..bb36927 --- /dev/null +++ b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/requestparameter/PersonalizeRequestParameterUtil.java @@ -0,0 +1,35 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.requestparameter; + +import org.opensearch.action.search.SearchRequest; +import org.opensearch.search.SearchExtBuilder; + +import java.util.List; +import java.util.stream.Collectors; + +public class PersonalizeRequestParameterUtil { + + public static PersonalizeRequestParameters getPersonalizeRequestParameters(SearchRequest searchRequest) { + PersonalizeRequestParametersExtBuilder personalizeRequestParameterExtBuilder = null; + if (searchRequest.source() != null && searchRequest.source().ext() != null && !searchRequest.source().ext().isEmpty()) { + List extBuilders = searchRequest.source().ext().stream() + .filter(extBuilder -> PersonalizeRequestParametersExtBuilder.NAME.equals(extBuilder.getWriteableName())) + .collect(Collectors.toList()); + + if (!extBuilders.isEmpty()) { + personalizeRequestParameterExtBuilder = (PersonalizeRequestParametersExtBuilder) extBuilders.get(0); + } + } + PersonalizeRequestParameters personalizeRequestParameters = null; + if (personalizeRequestParameterExtBuilder != null) { + personalizeRequestParameters = personalizeRequestParameterExtBuilder.getRequestParameters(); + } + return personalizeRequestParameters; + } +} diff --git a/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/requestparameter/PersonalizeRequestParameters.java b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/requestparameter/PersonalizeRequestParameters.java new file mode 100644 index 0000000..cb7a102 --- /dev/null +++ b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/requestparameter/PersonalizeRequestParameters.java @@ -0,0 +1,110 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.requestparameter; + +import org.opensearch.core.common.io.stream.StreamInput; +import org.opensearch.core.common.io.stream.StreamOutput; +import org.opensearch.core.common.io.stream.Writeable; +import org.opensearch.core.ParseField; +import org.opensearch.core.xcontent.ObjectParser; +import org.opensearch.core.xcontent.ToXContentObject; +import org.opensearch.core.xcontent.XContentBuilder; +import org.opensearch.core.xcontent.XContentParser; + +import java.io.IOException; +import java.util.Map; +import java.util.Objects; + +public class PersonalizeRequestParameters implements Writeable, ToXContentObject { + + static final String PERSONALIZE_REQUEST_PARAMETERS = "personalize_request_parameters"; + private static final String USER_ID_PARAMETER = "user_id"; + private static final String CONTEXT_PARAMETER = "context"; + + private static final ObjectParser PARSER; + private static final ParseField USER_ID = new ParseField(USER_ID_PARAMETER); + private static final ParseField CONTEXT = new ParseField(CONTEXT_PARAMETER); + + static { + PARSER = new ObjectParser<>(PERSONALIZE_REQUEST_PARAMETERS, PersonalizeRequestParameters::new); + PARSER.declareString(PersonalizeRequestParameters::setUserId, USER_ID); + PARSER.declareObject(PersonalizeRequestParameters::setContext,(XContentParser p, Void c) -> { + try { + return p.map(); + } catch (IOException e) { + throw new IllegalArgumentException("Error parsing Personalize context from request parameters", e); + } + }, CONTEXT); + } + + private String userId; + + private Map context; + + public PersonalizeRequestParameters() {} + + public PersonalizeRequestParameters(String userId, Map context) { + this.userId = userId; + this.context = context; + } + + public PersonalizeRequestParameters(StreamInput input) throws IOException { + this.userId = input.readString(); + this.context = input.readMap(); + } + + public String getUserId() { + return userId; + } + + public void setUserId(String userId) { + this.userId = userId; + } + + public Map getContext() { + return context; + } + + public void setContext(Map context) { + this.context = context; + } + + @Override + public void writeTo(StreamOutput out) throws IOException { + out.writeString(this.userId); + out.writeMap(this.context); + } + + @Override + public XContentBuilder toXContent(XContentBuilder builder, Params params) throws IOException { + builder.field(USER_ID.getPreferredName(), this.userId); + return builder.field(CONTEXT.getPreferredName(), this.context); + } + + public static PersonalizeRequestParameters parse(XContentParser parser) throws IOException { + PersonalizeRequestParameters requestParameters = PARSER.parse(parser, null); + return requestParameters; + } + + @Override + public boolean equals(Object o) { + if (this == o) return true; + if (o == null || getClass() != o.getClass()) return false; + + PersonalizeRequestParameters config = (PersonalizeRequestParameters) o; + + if (!userId.equals(config.userId)) return false; + if (context.size() != config.getContext().size()) return false; + return userId.equals(config.userId) && context.equals(config.getContext()); + } + + @Override + public int hashCode() { + return Objects.hash(userId, context); + } +} diff --git a/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/requestparameter/PersonalizeRequestParametersExtBuilder.java b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/requestparameter/PersonalizeRequestParametersExtBuilder.java new file mode 100644 index 0000000..a7b6e8c --- /dev/null +++ b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/requestparameter/PersonalizeRequestParametersExtBuilder.java @@ -0,0 +1,80 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.requestparameter; + +import org.apache.logging.log4j.LogManager; +import org.apache.logging.log4j.Logger; +import org.opensearch.core.common.io.stream.StreamInput; +import org.opensearch.core.common.io.stream.StreamOutput; +import org.opensearch.core.xcontent.XContentBuilder; +import org.opensearch.core.xcontent.XContentParser; +import org.opensearch.search.SearchExtBuilder; + +import java.io.IOException; +import java.util.Objects; + +import static org.opensearch.search.relevance.transformer.personalizeintelligentranking.requestparameter.PersonalizeRequestParameters.PERSONALIZE_REQUEST_PARAMETERS; + +public class PersonalizeRequestParametersExtBuilder extends SearchExtBuilder { + private static final Logger logger = LogManager.getLogger(PersonalizeRequestParametersExtBuilder.class); + public static final String NAME = PERSONALIZE_REQUEST_PARAMETERS; + private PersonalizeRequestParameters requestParameters; + + public PersonalizeRequestParametersExtBuilder() {} + + public PersonalizeRequestParametersExtBuilder(StreamInput input) throws IOException { + requestParameters = new PersonalizeRequestParameters(input); + } + + public PersonalizeRequestParameters getRequestParameters() { + return requestParameters; + } + + public void setRequestParameters(PersonalizeRequestParameters requestParameters) { + this.requestParameters = requestParameters; + } + + @Override + public int hashCode() { + return Objects.hash(this.getClass(), this.requestParameters); + } + + @Override + public boolean equals(Object obj) { + if (obj == null) { + return false; + } + if (!(obj instanceof PersonalizeRequestParametersExtBuilder)) { + return false; + } + PersonalizeRequestParametersExtBuilder o = (PersonalizeRequestParametersExtBuilder) obj; + return this.requestParameters.equals(o.requestParameters); + } + + @Override + public String getWriteableName() { + return NAME; + } + + @Override + public void writeTo(StreamOutput out) throws IOException { + requestParameters.writeTo(out); + } + + @Override + public XContentBuilder toXContent(XContentBuilder builder, Params params) throws IOException { + return builder.value(requestParameters); + } + + public static PersonalizeRequestParametersExtBuilder parse(XContentParser parser) throws IOException{ + PersonalizeRequestParametersExtBuilder extBuilder = new PersonalizeRequestParametersExtBuilder(); + PersonalizeRequestParameters requestParameters = PersonalizeRequestParameters.parse(parser); + extBuilder.setRequestParameters(requestParameters); + return extBuilder; + } +} diff --git a/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/reranker/PersonalizedRanker.java b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/reranker/PersonalizedRanker.java new file mode 100644 index 0000000..8470f9f --- /dev/null +++ b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/reranker/PersonalizedRanker.java @@ -0,0 +1,22 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.reranker; + +import org.opensearch.search.SearchHits; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.requestparameter.PersonalizeRequestParameters; + +public interface PersonalizedRanker { + + /** + * Re rank search hits + * @param hits Search hits to re rank + * @param requestParameters Request parameters for Personalize present in search request + * @return Re ranked search hits + */ + SearchHits rerank(SearchHits hits, PersonalizeRequestParameters requestParameters); +} diff --git a/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/reranker/PersonalizedRankerFactory.java b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/reranker/PersonalizedRankerFactory.java new file mode 100644 index 0000000..ead8dd8 --- /dev/null +++ b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/reranker/PersonalizedRankerFactory.java @@ -0,0 +1,43 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.reranker; + +import org.apache.logging.log4j.LogManager; +import org.apache.logging.log4j.Logger; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.client.PersonalizeClient; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration.PersonalizeIntelligentRankerConfiguration; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.reranker.impl.AmazonPersonalizedRankerImpl; + +import static org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration.Constants.AMAZON_PERSONALIZED_RANKING_RECIPE_NAME; +import static org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration.Constants.AMAZON_PERSONALIZED_RANKING_V2_RECIPE_NAME; + +/** + * Factory for creating Personalize ranker instance based on Personalize ranker configuration + */ +public class PersonalizedRankerFactory { + private static final Logger logger = LogManager.getLogger(PersonalizedRankerFactory.class); + + /** + * Create an instance of Personalize ranker based on ranker configuration + * @param config Personalize ranker configuration + * @param client Personalize client + * @return Personalize ranker instance + */ + public PersonalizedRanker getPersonalizedRanker(PersonalizeIntelligentRankerConfiguration config, PersonalizeClient client){ + PersonalizedRanker ranker = null; + String recipeInConfig = config.getRecipe(); + if (recipeInConfig.equals(AMAZON_PERSONALIZED_RANKING_RECIPE_NAME) + || recipeInConfig.equals(AMAZON_PERSONALIZED_RANKING_V2_RECIPE_NAME)) { + ranker = new AmazonPersonalizedRankerImpl(config, client); + } else { + logger.error("Personalize recipe provided in configuration is not supported for re ranking search results"); + //TODO : throw user error exception + } + return ranker; + } +} diff --git a/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/reranker/impl/AmazonPersonalizedRankerImpl.java b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/reranker/impl/AmazonPersonalizedRankerImpl.java new file mode 100644 index 0000000..142dceb --- /dev/null +++ b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/reranker/impl/AmazonPersonalizedRankerImpl.java @@ -0,0 +1,171 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.reranker.impl; + +import com.amazonaws.AmazonServiceException; +import com.amazonaws.services.personalizeruntime.model.GetPersonalizedRankingRequest; +import com.amazonaws.services.personalizeruntime.model.GetPersonalizedRankingResult; +import com.amazonaws.services.personalizeruntime.model.PredictedItem; +import org.apache.logging.log4j.LogManager; +import org.apache.logging.log4j.Logger; +import org.opensearch.ingest.ConfigurationUtils; +import org.opensearch.search.SearchHit; +import org.opensearch.search.SearchHits; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.PersonalizeRankingResponseProcessor; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.client.PersonalizeClient; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration.PersonalizeIntelligentRankerConfiguration; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.requestparameter.PersonalizeRequestParameters; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.reranker.PersonalizedRanker; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.utils.ValidationUtil; + +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Comparator; +import java.util.List; +import java.util.LinkedList; +import java.util.Map; +import java.util.stream.Collectors; + +/** + * Personalize Re Ranker implementation using Amazon Personalized Ranking recipe + */ +public class AmazonPersonalizedRankerImpl implements PersonalizedRanker { + private static final Logger logger = LogManager.getLogger(AmazonPersonalizedRankerImpl.class); + private final PersonalizeIntelligentRankerConfiguration rankerConfig; + private final PersonalizeClient personalizeClient; + + public AmazonPersonalizedRankerImpl(PersonalizeIntelligentRankerConfiguration config, + PersonalizeClient client) { + this.rankerConfig = config; + this.personalizeClient = client; + } + + /** + * Re rank search hits using Personalize campaign that uses Personalized Ranking recipe + * @param hits search hits returned by open search + * @param requestParameters request parameters for Personalize present in search request + * @return search hots re ranked using Amazon Personalize + */ + @Override + public SearchHits rerank(SearchHits hits, PersonalizeRequestParameters requestParameters) { + try { + validatePersonalizeRequestParams(requestParameters); + List originalHits = Arrays.asList(hits.getHits()); + // Do not make Personalize call if weight is zero which implies Personalization is turned off. + if (rankerConfig.getWeight() == 0) { + logger.info("Not applying Personalized ranking. Given value for weight configuration: {}", rankerConfig.getWeight()); + return hits; + } + String itemIdfield = rankerConfig.getItemIdField(); + List documentIdsToRank; + // If item field is not specified in the configuration then use default _id field. + if (itemIdfield != null && !itemIdfield.isBlank()) { + documentIdsToRank = originalHits.stream() + .filter(h -> h.getSourceAsMap().get(itemIdfield) != null) + .map(h -> h.getSourceAsMap().get(itemIdfield).toString()) + .collect(Collectors.toList()); + } else { + documentIdsToRank = originalHits.stream() + .filter(h -> h.getId() != null) + .map(h -> h.getId()) + .collect(Collectors.toList()); + } + if (documentIdsToRank.size() == 0) { + throw ConfigurationUtils.newConfigurationException(PersonalizeRankingResponseProcessor.TYPE, "", "item_id_field", + "no item ids found to apply Personalized reranking. Please check configured value for item_id_field"); + } + logger.info("Document Ids to re-rank with Personalize: {}", Arrays.toString(documentIdsToRank.toArray())); + String userId = requestParameters.getUserId(); + Map context = requestParameters.getContext() != null ? + requestParameters.getContext().entrySet().stream() + .collect(Collectors.toMap(Map.Entry::getKey, e -> (String)e.getValue())) + : null; + logger.info("User ID from personalize request parameters - User ID: {}", userId); + if (context != null && !context.isEmpty()) { + logger.info("Personalize context provided in the search request"); + } + + GetPersonalizedRankingRequest personalizeRequest = new GetPersonalizedRankingRequest() + .withCampaignArn(rankerConfig.getPersonalizeCampaign()) + .withInputList(documentIdsToRank) + .withContext(context) + .withUserId(userId); + GetPersonalizedRankingResult result = personalizeClient.getPersonalizedRanking(personalizeRequest); + + SearchHits personalizedHits = combineScores(hits, result); + return personalizedHits; + } catch (AmazonServiceException e) { + logger.error("Exception while calling personalize campaign: {}", e.getMessage()); + int statusCode = e.getStatusCode(); + if (ValidationUtil.is4xxError(statusCode)) { + throw new IllegalArgumentException(e); + } + throw e; + } + catch (Exception ex) { + logger.error("Failed to re rank with Personalize.", ex); + throw ex; + } + } + + //Combine open search hits and personalize campaign response + private SearchHits combineScores(SearchHits originalHits, GetPersonalizedRankingResult personalizedRankingResult) { + List personalziedRanking = personalizedRankingResult.getPersonalizedRanking(); + List personalizedRankedItemsList = new LinkedList<>(); + for (PredictedItem item : personalziedRanking) { + personalizedRankedItemsList.add(item.getItemId()); + } + int totalHits = originalHits.getHits().length; + List rerankedHits = new ArrayList<>(totalHits); + float maxScore = 0f; + double weight = rankerConfig.getWeight(); + for (int i = 0 ; i < totalHits ; i++) { + String openSearchItemId; + SearchHit hit = originalHits.getAt(i); + String itemIdField = rankerConfig.getItemIdField(); + if (itemIdField != null && !(itemIdField.isBlank())) { + openSearchItemId = hit.getSourceAsMap().get(rankerConfig.getItemIdField()).toString(); + } else { + openSearchItemId = hit.getId(); + } + int openSearchRank = i + 1; + int personalizedRank = personalizedRankedItemsList.indexOf(openSearchItemId) + 1; + float combinedScore = (float) (((1- weight) / (Math.log(openSearchRank + 1) / Math.log(2))) + + ((weight) / (Math.log(personalizedRank + 1) / Math.log(2)))); + maxScore = Math.max(maxScore, combinedScore); + hit.score(combinedScore); + rerankedHits.add(hit); + } + rerankedHits.sort(Comparator.comparing(SearchHit::getScore).reversed()); + return new SearchHits(rerankedHits.toArray(new SearchHit[0]), originalHits.getTotalHits(), maxScore); + } + + /** + * Validate Personalize configuration for calling Personalize service + * @param requestParameters Request parameters for Personalize present in search request + */ + private void validatePersonalizeRequestParams(PersonalizeRequestParameters requestParameters) { + if (requestParameters == null || requestParameters.getUserId() == null || requestParameters.getUserId().isBlank()) { + throw ConfigurationUtils.newConfigurationException(PersonalizeRankingResponseProcessor.TYPE, "", "user_id", + "required Personalize request parameter is missing"); + } + if (requestParameters.getContext() != null) { + try { + requestParameters.getContext().entrySet().stream().forEach(e -> isValidPersonalizeContext(e)); + } catch (IllegalArgumentException iae) { + throw ConfigurationUtils.newConfigurationException(PersonalizeRankingResponseProcessor.TYPE, "", "context", iae.getMessage()); + } + } + } + + private void isValidPersonalizeContext(Map.Entry contextEntry) throws IllegalArgumentException { + if (!(contextEntry.getValue() instanceof String)) { + throw new IllegalArgumentException("Personalize context value is not of type String. Invalid context value: " + contextEntry.getValue()); + } + } +} diff --git a/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/utils/ValidationUtil.java b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/utils/ValidationUtil.java new file mode 100644 index 0000000..c6c92f7 --- /dev/null +++ b/amazon-personalize-ranking/src/main/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/utils/ValidationUtil.java @@ -0,0 +1,71 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.utils; +import com.amazonaws.arn.Arn; +import org.opensearch.ingest.ConfigurationUtils; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration.PersonalizeIntelligentRankerConfiguration; + +import java.util.Arrays; +import java.util.Set; +import java.util.HashSet; + +import static org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration.Constants.AMAZON_PERSONALIZED_RANKING_RECIPE_NAME; +import static org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration.Constants.AMAZON_PERSONALIZED_RANKING_V2_RECIPE_NAME; + +public class ValidationUtil { + private static Set SUPPORTED_PERSONALIZE_RECIPES = new HashSet<>(Arrays.asList( + AMAZON_PERSONALIZED_RANKING_RECIPE_NAME, + AMAZON_PERSONALIZED_RANKING_V2_RECIPE_NAME + )); + + /** + * Validate Personalize configuration for calling Personalize service. + * Throws OpenSearchParseException type exception if validation fails. + * @param config Personalize intelligent ranker configuration + * @param processorType Name of search pipeline processor + * @param processorTag Name of processor tag + */ + public static void validatePersonalizeIntelligentRankerConfiguration (PersonalizeIntelligentRankerConfiguration config, + String processorType, + String processorTag + ) { + // Validate weight value + if (config.getWeight() < 0.0 || config.getWeight() > 1.0) { + throw ConfigurationUtils.newConfigurationException(processorType, processorTag, "weight", "invalid value for weight"); + } + // Validate Personalize campaign ARN + if(!isValidCampaignOrRoleArn(config.getPersonalizeCampaign(), "personalize")) { + throw ConfigurationUtils.newConfigurationException(processorType, processorTag, "campaign_arn", "invalid format for Personalize campaign arn"); + } + // Validate IAM Role Arn for Personalize access + String iamRoleArn = config.getIamRoleArn(); + if(!(iamRoleArn == null || iamRoleArn.isBlank()) && !isValidCampaignOrRoleArn(iamRoleArn, "iam")) { + throw ConfigurationUtils.newConfigurationException(processorType, processorTag, "iam_role_arn", "invalid format for Personalize iam role arn"); + } + // Validate Personalize recipe + if(!SUPPORTED_PERSONALIZE_RECIPES.contains(config.getRecipe())) { + throw ConfigurationUtils.newConfigurationException(processorType, processorTag, "recipe", "not supported recipe provided"); + } + } + + private static boolean isValidCampaignOrRoleArn(String arn, String expectedService) { + try { + Arn arnObj = Arn.fromString(arn); + return arnObj.getService().equals(expectedService); + } catch (IllegalArgumentException iae) { + return false; + } + } + + public static boolean is4xxError(int statusCode){ + if (statusCode >= 400 && statusCode < 500) { + return true; + } + return false; + } +} diff --git a/amazon-personalize-ranking/src/main/plugin-metadata/plugin-security.policy b/amazon-personalize-ranking/src/main/plugin-metadata/plugin-security.policy new file mode 100644 index 0000000..b16dfe0 --- /dev/null +++ b/amazon-personalize-ranking/src/main/plugin-metadata/plugin-security.policy @@ -0,0 +1,15 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ + +grant { + permission java.lang.RuntimePermission "accessDeclaredMembers"; + permission java.lang.reflect.ReflectPermission "suppressAccessChecks"; + + permission java.net.SocketPermission "*", "connect,resolve"; + permission java.lang.RuntimePermission "getClassLoader"; +}; diff --git a/src/test/java/org/opensearch/search/relevance/SearchRelevancePluginIT.java b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/AmazonPersonalizeRankingPluginIT.java similarity index 85% rename from src/test/java/org/opensearch/search/relevance/SearchRelevancePluginIT.java rename to amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/AmazonPersonalizeRankingPluginIT.java index 63ebfdc..f72f947 100644 --- a/src/test/java/org/opensearch/search/relevance/SearchRelevancePluginIT.java +++ b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/AmazonPersonalizeRankingPluginIT.java @@ -15,7 +15,7 @@ import java.io.IOException; -public class SearchRelevancePluginIT extends OpenSearchRestTestCase { +public class AmazonPersonalizeRankingPluginIT extends OpenSearchRestTestCase { public void testPluginInstalled() throws IOException, ParseException { Response response = client().performRequest(new Request("GET", "/_cat/plugins")); @@ -23,6 +23,6 @@ public void testPluginInstalled() throws IOException, ParseException { logger.info("response body: {}", body); assertNotNull(body); - assertTrue(body.contains("search-processor")); + assertTrue(body.contains("amazon-personalize-ranking")); } } diff --git a/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/PersonalizeRankingResponseProcessorTests.java b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/PersonalizeRankingResponseProcessorTests.java new file mode 100644 index 0000000..f83bbcb --- /dev/null +++ b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/PersonalizeRankingResponseProcessorTests.java @@ -0,0 +1,435 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking; + +import com.amazonaws.http.IdleConnectionReaper; +import org.apache.lucene.search.TotalHits; +import org.opensearch.OpenSearchParseException; +import org.opensearch.action.search.SearchRequest; +import org.opensearch.action.search.SearchResponse; +import org.opensearch.action.search.SearchResponseSections; +import org.opensearch.action.search.ShardSearchFailure; +import org.opensearch.common.settings.Settings; +import org.opensearch.env.Environment; +import org.opensearch.env.TestEnvironment; +import org.opensearch.search.SearchHit; +import org.opensearch.search.SearchHits; +import org.opensearch.search.pipeline.Processor; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.client.PersonalizeClient; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.client.PersonalizeClientSettings; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.requestparameter.PersonalizeRequestParameters; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.utils.PersonalizeRuntimeTestUtil; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.utils.SearchTestUtil; +import org.opensearch.test.OpenSearchTestCase; + +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.Objects; +import java.util.stream.Collectors; + +import static org.mockito.Mockito.mock; +import static org.opensearch.search.relevance.transformer.personalizeintelligentranking.PersonalizeRankingResponseProcessor.TYPE; +import static org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration.Constants.AMAZON_PERSONALIZED_RANKING_RECIPE_NAME; +import static org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration.Constants.AMAZON_PERSONALIZED_RANKING_V2_RECIPE_NAME; + +public class PersonalizeRankingResponseProcessorTests extends OpenSearchTestCase { + + private static final Processor.PipelineContext UPDATE_CONTEXT = new Processor.PipelineContext(Processor.PipelineSource.UPDATE_PIPELINE); + private static final Processor.PipelineContext VALIDATE_CONTEXT = new Processor.PipelineContext(Processor.PipelineSource.VALIDATE_PIPELINE); + private final Settings settings = buildEnvSettings(Settings.EMPTY); + private final Environment env = TestEnvironment.newEnvironment(settings); + private static final String PERSONALIZE_CAMPAIGN = "arn:aws:personalize:us-west-2:000000000000:campaign/test-campaign"; + private static final String IAM_ROLE_ARN = "arn:aws:iam::000000000000:role/test"; + private static final String ITEM_ID_FIELD = "ITEM_ID"; + private static final String REGION = "us-west-2"; + private static final double WEIGHT = 1.0; + private static final int NUM_HITS = 10; + + private final PersonalizeClientSettings clientSettings = PersonalizeClientSettings.getClientSettings(env.settings()); + + public void testCreateFactoryThrowsExceptionWithEmptyConfig() { + PersonalizeRankingResponseProcessor.Factory factory + = new PersonalizeRankingResponseProcessor.Factory(this.clientSettings); + expectThrows(OpenSearchParseException.class, () -> factory.create( + Collections.emptyMap(), + null, + null, + false, + Collections.emptyMap(), + UPDATE_CONTEXT + )); + IdleConnectionReaper.shutdown(); + } + + public void testFactoryValidations() { + PersonalizeRankingResponseProcessor.Factory factory + = new PersonalizeRankingResponseProcessor.Factory(this.clientSettings); + // Test config without campaign + Map configuration = new HashMap<>(); + configuration.put("item_id_field", ITEM_ID_FIELD); + configuration.put("recipe", AMAZON_PERSONALIZED_RANKING_RECIPE_NAME); + configuration.put("weight", String.valueOf(WEIGHT)); + configuration.put("iam_role_arn", IAM_ROLE_ARN); + configuration.put("aws_region", REGION); + + expectThrows(OpenSearchParseException.class, () -> factory.create( + Collections.emptyMap(), + null, + null, + false, + configuration, + VALIDATE_CONTEXT + )); + configuration.clear(); + + // Test config without recipe + configuration.put("campaign_arn", PERSONALIZE_CAMPAIGN); + configuration.put("item_id_field", ITEM_ID_FIELD); + configuration.put("weight", String.valueOf(WEIGHT)); + configuration.put("iam_role_arn", IAM_ROLE_ARN); + configuration.put("aws_region", REGION); + + expectThrows(OpenSearchParseException.class, () -> factory.create( + Collections.emptyMap(), + null, + null, + false, + configuration, + VALIDATE_CONTEXT + )); + configuration.clear(); + + // Test config without region + configuration.put("campaign_arn", PERSONALIZE_CAMPAIGN); + configuration.put("item_id_field", ITEM_ID_FIELD); + configuration.put("recipe", AMAZON_PERSONALIZED_RANKING_RECIPE_NAME); + configuration.put("weight", String.valueOf(WEIGHT)); + configuration.put("iam_role_arn", IAM_ROLE_ARN); + expectThrows(OpenSearchParseException.class, () -> factory.create( + Collections.emptyMap(), + null, + null, + false, + configuration, + VALIDATE_CONTEXT + )); + configuration.clear(); + + // Test config without weight + configuration.put("campaign_arn", PERSONALIZE_CAMPAIGN); + configuration.put("item_id_field", ITEM_ID_FIELD); + configuration.put("recipe", AMAZON_PERSONALIZED_RANKING_RECIPE_NAME); + configuration.put("iam_role_arn", IAM_ROLE_ARN); + configuration.put("aws_region", REGION); + expectThrows(OpenSearchParseException.class, () -> factory.create( + Collections.emptyMap(), + null, + null, + false, + configuration, + VALIDATE_CONTEXT + )); + configuration.clear(); + + // Test configuration with invalid weight value + configuration.put("campaign_arn", PERSONALIZE_CAMPAIGN); + configuration.put("item_id_field", ITEM_ID_FIELD); + configuration.put("recipe", AMAZON_PERSONALIZED_RANKING_RECIPE_NAME); + configuration.put("weight", "invalid"); + configuration.put("iam_role_arn", IAM_ROLE_ARN); + configuration.put("aws_region", REGION); + expectThrows(OpenSearchParseException.class, () -> factory.create( + Collections.emptyMap(), + null, + null, + false, + configuration, + VALIDATE_CONTEXT + )); + configuration.clear(); + + configuration.put("campaign_arn", PERSONALIZE_CAMPAIGN); + configuration.put("item_id_field", ITEM_ID_FIELD); + configuration.put("recipe", AMAZON_PERSONALIZED_RANKING_RECIPE_NAME); + configuration.put("weight", String.valueOf(WEIGHT)); + configuration.put("iam_role_arn", IAM_ROLE_ARN); + configuration.put("aws_region", REGION); + + // Test that we don't create client on validation + configuration.putAll(buildPersonalizeResponseProcessorConfig()); + PersonalizeRankingResponseProcessor processor = factory.create(Collections.emptyMap(), null, null, false, configuration, VALIDATE_CONTEXT); + assertNull(processor.getPersonalizeClient()); + + // Test that we fail on valid configuration in search request context + expectThrows(IllegalStateException.class, () -> factory.create( + Collections.emptyMap(), + null, + null, + false, + buildPersonalizeResponseProcessorConfig(), + new Processor.PipelineContext(Processor.PipelineSource.SEARCH_REQUEST))); + } + + public void testCreateFactoryWithAllPersonalizeConfig() throws Exception { + PersonalizeRankingResponseProcessor.Factory factory + = new PersonalizeRankingResponseProcessor.Factory(this.clientSettings); + + Map configuration = buildPersonalizeResponseProcessorConfig(); + + PersonalizeRankingResponseProcessor personalizeResponseProcessor = + factory.create(Collections.emptyMap(), "testTag", "testingAllFields", false, configuration, UPDATE_CONTEXT); + + assertEquals(TYPE, personalizeResponseProcessor.getType()); + assertEquals("testTag", personalizeResponseProcessor.getTag()); + assertEquals("testingAllFields", personalizeResponseProcessor.getDescription()); + IdleConnectionReaper.shutdown(); + } + + public void testProcessorWithNoHits() throws Exception { + PersonalizeClient mockClient = mock(PersonalizeClient.class); + PersonalizeRankingResponseProcessor.Factory factory + = new PersonalizeRankingResponseProcessor.Factory(this.clientSettings, (cp, r) -> mockClient); + + Map configuration = buildPersonalizeResponseProcessorConfig(); + + PersonalizeRankingResponseProcessor personalizeResponseProcessor = + factory.create(Collections.emptyMap(), "testTag", "testingAllFields", false, configuration, UPDATE_CONTEXT); + SearchRequest searchRequest = new SearchRequest(); + SearchHits hits = new SearchHits(new SearchHit[0], new TotalHits(0, TotalHits.Relation.EQUAL_TO), 0.0f); + SearchResponseSections searchResponseSections = new SearchResponseSections(hits, null, null, false, false, null, 0); + SearchResponse searchResponse = new SearchResponse(searchResponseSections, null, 1, 1, 0, 1, new ShardSearchFailure[0], null); + + SearchResponse response = personalizeResponseProcessor.processResponse(searchRequest, searchResponse); + assertEquals(hits.getTotalHits().value, response.getHits().getTotalHits().value); + IdleConnectionReaper.shutdown(); + } + + public void testProcessorWithPersonalizeContext() throws Exception { + PersonalizeClient mockClient = PersonalizeRuntimeTestUtil.buildMockPersonalizeClient(); + + PersonalizeRankingResponseProcessor.Factory factory + = new PersonalizeRankingResponseProcessor.Factory(this.clientSettings, (cp, r) -> mockClient); + + Map configuration = buildPersonalizeResponseProcessorConfig(); + PersonalizeRankingResponseProcessor personalizeResponseProcessor = + factory.create(Collections.emptyMap(), "testTag", "testingAllFields", false, configuration, UPDATE_CONTEXT); + + Map personalizeContext = new HashMap<>(); + personalizeContext.put("contextKey2", "contextValue2"); + + SearchResponse personalizedResponse = + createPersonalizedRankingProcessorResponse(personalizeResponseProcessor, personalizeContext, NUM_HITS); + + List transformedHits = Arrays.asList(personalizedResponse.getHits().getHits()); + List rerankedDocumentIds; + rerankedDocumentIds = transformedHits.stream() + .filter(h -> h.getSourceAsMap().get(ITEM_ID_FIELD) != null) + .map(h -> h.getSourceAsMap().get(ITEM_ID_FIELD).toString()) + .collect(Collectors.toList()); + + ArrayList expectedRankedDocumentIds = PersonalizeRuntimeTestUtil.expectedRankedItemIdsForGivenWeight(NUM_HITS, 1); + assertEquals(expectedRankedDocumentIds, rerankedDocumentIds); + IdleConnectionReaper.shutdown(); + } + + public void testProcessorWithHitsWithInvalidPersonalizeContext() throws Exception { + PersonalizeClient mockClient = PersonalizeRuntimeTestUtil.buildMockPersonalizeClient();; + + PersonalizeRankingResponseProcessor.Factory factory + = new PersonalizeRankingResponseProcessor.Factory(this.clientSettings, (cp, r) -> mockClient); + + Map configuration = buildPersonalizeResponseProcessorConfig(); + PersonalizeRankingResponseProcessor personalizeResponseProcessor = + factory.create(Collections.emptyMap(), "testTag", "testingAllFields", false, configuration, UPDATE_CONTEXT); + + Map personalizeContext = new HashMap<>(); + personalizeContext.put("contextKey2", 5); + + expectThrows(OpenSearchParseException.class, () -> + createPersonalizedRankingProcessorResponse(personalizeResponseProcessor, personalizeContext, NUM_HITS)); + IdleConnectionReaper.shutdown(); + } + + public void testPersonalizeRankingResponse() throws Exception { + PersonalizeClient personalizeClient = PersonalizeRuntimeTestUtil.buildMockPersonalizeClient(); + + PersonalizeRankingResponseProcessor.Factory factory + = new PersonalizeRankingResponseProcessor.Factory(this.clientSettings, (cp, r) -> personalizeClient); + + String itemField = "ITEM_ID"; + Map configuration = buildPersonalizeResponseProcessorConfig(); + + PersonalizeRankingResponseProcessor responseProcessor = + factory.create(Collections.emptyMap(), "testTag", "testingAllFields", false, configuration, UPDATE_CONTEXT); + + SearchResponse personalizedResponse = createPersonalizedRankingProcessorResponse(responseProcessor, null, NUM_HITS); + + List transformedHits = Arrays.asList(personalizedResponse.getHits().getHits()); + List rerankedDocumentIds; + rerankedDocumentIds = transformedHits.stream() + .filter(h -> h.getSourceAsMap().get(itemField) != null) + .map(h -> h.getSourceAsMap().get(itemField).toString()) + .collect(Collectors.toList()); + + ArrayList expectedRankedDocumentIds = PersonalizeRuntimeTestUtil.expectedRankedItemIdsForGivenWeight(NUM_HITS, 1); + assertEquals(expectedRankedDocumentIds, rerankedDocumentIds); + IdleConnectionReaper.shutdown(); + } + + public void testPersonalizeRankingV2Response() throws Exception { + PersonalizeClient personalizeClient = PersonalizeRuntimeTestUtil.buildMockPersonalizeClient(); + + PersonalizeRankingResponseProcessor.Factory factory + = new PersonalizeRankingResponseProcessor.Factory(this.clientSettings, (cp, r) -> personalizeClient); + + String itemField = "ITEM_ID"; + Map configuration = buildPersonalizeResponseProcessorConfig(); + configuration.put("recipe", AMAZON_PERSONALIZED_RANKING_V2_RECIPE_NAME); + + PersonalizeRankingResponseProcessor responseProcessor = + factory.create(Collections.emptyMap(), "testTag", "testingAllFields", false, configuration, UPDATE_CONTEXT); + + SearchResponse personalizedResponse = createPersonalizedRankingProcessorResponse(responseProcessor, null, NUM_HITS); + + List transformedHits = Arrays.asList(personalizedResponse.getHits().getHits()); + List rerankedDocumentIds; + rerankedDocumentIds = transformedHits.stream() + .filter(h -> h.getSourceAsMap().get(itemField) != null) + .map(h -> h.getSourceAsMap().get(itemField).toString()) + .collect(Collectors.toList()); + + ArrayList expectedRankedDocumentIds = PersonalizeRuntimeTestUtil.expectedRankedItemIdsForGivenWeight(NUM_HITS, 1); + assertEquals(expectedRankedDocumentIds, rerankedDocumentIds); + IdleConnectionReaper.shutdown(); + } + + public void testPersonalizeRankingV2ResponseWithInvalidItemIdFieldName() throws Exception { + PersonalizeClient personalizeClient = PersonalizeRuntimeTestUtil.buildMockPersonalizeClient(); + + PersonalizeRankingResponseProcessor.Factory factory + = new PersonalizeRankingResponseProcessor.Factory(this.clientSettings, (cp, r) -> personalizeClient); + + String itemFieldInvalid = "ITEM_ID_NOT_VALID"; + Map configuration = buildPersonalizeResponseProcessorConfig(); + configuration.put("recipe", AMAZON_PERSONALIZED_RANKING_V2_RECIPE_NAME); + configuration.put("item_id_field", itemFieldInvalid); + + PersonalizeRankingResponseProcessor responseProcessor = + factory.create(Collections.emptyMap(), "testTag", "testingAllFields", false, configuration, UPDATE_CONTEXT); + + expectThrows(OpenSearchParseException.class, () -> + createPersonalizedRankingProcessorResponse(responseProcessor, null, NUM_HITS)); + IdleConnectionReaper.shutdown(); + } + + public void testPersonalizeRankingV2ResponseWithDefaultItemIdField() throws Exception { + PersonalizeClient personalizeClient = PersonalizeRuntimeTestUtil.buildMockPersonalizeClient(); + + PersonalizeRankingResponseProcessor.Factory factory + = new PersonalizeRankingResponseProcessor.Factory(this.clientSettings, (cp, r) -> personalizeClient); + + String itemIdFieldEmpty = ""; + Map configuration = buildPersonalizeResponseProcessorConfig(); + configuration.put("item_id_field", itemIdFieldEmpty); + configuration.put("recipe", AMAZON_PERSONALIZED_RANKING_V2_RECIPE_NAME); + + PersonalizeRankingResponseProcessor responseProcessor = + factory.create(Collections.emptyMap(), "testTag", "testingAllFields", false, configuration, UPDATE_CONTEXT); + + SearchResponse personalizedResponse = createPersonalizedRankingProcessorResponse(responseProcessor, null, NUM_HITS); + + List transformedHits = Arrays.asList(personalizedResponse.getHits().getHits()); + List rerankedDocumentIds; + rerankedDocumentIds = transformedHits.stream() + .map(SearchHit::getId) + .filter(Objects::nonNull) + .collect(Collectors.toList()); + + ArrayList expectedRankedDocumentIds = PersonalizeRuntimeTestUtil.expectedRankedItemIdsForGivenWeight(NUM_HITS, 1); + + assertEquals(expectedRankedDocumentIds, rerankedDocumentIds); + IdleConnectionReaper.shutdown(); + } + + public void testPersonalizeRankingResponseWithInvalidItemIdFieldName() throws Exception { + PersonalizeClient personalizeClient = PersonalizeRuntimeTestUtil.buildMockPersonalizeClient(); + + PersonalizeRankingResponseProcessor.Factory factory + = new PersonalizeRankingResponseProcessor.Factory(this.clientSettings, (cp, r) -> personalizeClient); + + String itemFieldInvalid = "ITEM_ID_NOT_VALID"; + Map configuration = buildPersonalizeResponseProcessorConfig(); + configuration.put("item_id_field", itemFieldInvalid); + + PersonalizeRankingResponseProcessor responseProcessor = + factory.create(Collections.emptyMap(), "testTag", "testingAllFields", false, configuration, UPDATE_CONTEXT); + + expectThrows(OpenSearchParseException.class, () -> + createPersonalizedRankingProcessorResponse(responseProcessor, null, NUM_HITS)); + IdleConnectionReaper.shutdown(); + } + + public void testPersonalizeRankingResponseWithDefaultItemIdField() throws Exception { + PersonalizeClient personalizeClient = PersonalizeRuntimeTestUtil.buildMockPersonalizeClient(); + + PersonalizeRankingResponseProcessor.Factory factory + = new PersonalizeRankingResponseProcessor.Factory(this.clientSettings, (cp, r) -> personalizeClient); + + String itemIdFieldEmpty = ""; + Map configuration = buildPersonalizeResponseProcessorConfig(); + configuration.put("item_id_field", itemIdFieldEmpty); + + PersonalizeRankingResponseProcessor responseProcessor = + factory.create(Collections.emptyMap(), "testTag", "testingAllFields", false, configuration, UPDATE_CONTEXT); + + SearchResponse personalizedResponse = createPersonalizedRankingProcessorResponse(responseProcessor, null, NUM_HITS); + + List transformedHits = Arrays.asList(personalizedResponse.getHits().getHits()); + List rerankedDocumentIds; + rerankedDocumentIds = transformedHits.stream() + .map(SearchHit::getId) + .filter(Objects::nonNull) + .collect(Collectors.toList()); + + ArrayList expectedRankedDocumentIds = PersonalizeRuntimeTestUtil.expectedRankedItemIdsForGivenWeight(NUM_HITS, 1); + + assertEquals(expectedRankedDocumentIds, rerankedDocumentIds); + IdleConnectionReaper.shutdown(); + } + + private SearchResponse createPersonalizedRankingProcessorResponse(PersonalizeRankingResponseProcessor responseProcessor, + Map personalizeContext, + int numHits) throws Exception { + + PersonalizeRequestParameters personalizeRequestParams = new PersonalizeRequestParameters("user_1", personalizeContext); + SearchRequest request = SearchTestUtil.createSearchRequestWithPersonalizeRequest(personalizeRequestParams); + + SearchHits searchHits = SearchTestUtil.getSampleSearchHitsForPersonalize(numHits); + SearchResponseSections searchResponseSections = new SearchResponseSections(searchHits, null, null, false, false, null, 0); + SearchResponse searchResponse = new SearchResponse(searchResponseSections, null, 1, 1, 0, 1, new ShardSearchFailure[0], null); + + SearchResponse personalizedResponse = responseProcessor.processResponse(request, searchResponse); + + return personalizedResponse; + } + + private Map buildPersonalizeResponseProcessorConfig() { + Map configuration = new HashMap<>(); + configuration.put("campaign_arn", PERSONALIZE_CAMPAIGN); + configuration.put("item_id_field", ITEM_ID_FIELD); + configuration.put("recipe", AMAZON_PERSONALIZED_RANKING_RECIPE_NAME); + configuration.put("weight", String.valueOf(WEIGHT)); + configuration.put("iam_role_arn", IAM_ROLE_ARN); + configuration.put("aws_region", REGION); + return configuration; + } +} diff --git a/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeClientSettingsTests.java b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeClientSettingsTests.java new file mode 100644 index 0000000..bf1def8 --- /dev/null +++ b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeClientSettingsTests.java @@ -0,0 +1,77 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.client; + +import com.amazonaws.auth.AWSCredentials; +import com.amazonaws.auth.AWSSessionCredentials; +import org.opensearch.common.settings.SecureSetting; +import org.opensearch.common.settings.Setting; +import org.opensearch.common.settings.SettingsException; +import org.opensearch.core.common.settings.SecureString; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.utils.PersonalizeClientSettingsTestUtil; +import org.opensearch.test.OpenSearchTestCase; + +import java.io.IOException; +import java.util.Arrays; +import java.util.Collection; +import java.util.List; + +import static org.opensearch.search.relevance.transformer.personalizeintelligentranking.utils.PersonalizeClientSettingsTestUtil.ACCESS_KEY; +import static org.opensearch.search.relevance.transformer.personalizeintelligentranking.utils.PersonalizeClientSettingsTestUtil.SECRET_KEY; +import static org.opensearch.search.relevance.transformer.personalizeintelligentranking.utils.PersonalizeClientSettingsTestUtil.SESSION_TOKEN; + +public class PersonalizeClientSettingsTests extends OpenSearchTestCase { + + public void testWithBasicCredentials() throws IOException { + PersonalizeClientSettings clientSettings = PersonalizeClientSettingsTestUtil.buildClientSettings(true, true, false); + AWSCredentials credentials = clientSettings.getCredentials(); + assertEquals(ACCESS_KEY, credentials.getAWSAccessKeyId()); + assertEquals(SECRET_KEY, credentials.getAWSSecretKey()); + assertFalse(credentials instanceof AWSSessionCredentials); + } + + public void testWithGetAllSetting() throws IOException { + PersonalizeClientSettings clientSettings = PersonalizeClientSettingsTestUtil.buildClientSettings(true, true, true); + assertEquals(clientSettings.getAllSettings().size(), 3); + Setting ACCESS_KEY_SETTING = SecureSetting.secureString("personalized_search_ranking.aws.access_key", null); + Setting SECRET_KEY_SETTING = SecureSetting.secureString("personalized_search_ranking.aws.secret_key", null); + Setting SESSION_TOKEN_SETTING = SecureSetting.secureString("personalized_search_ranking.aws.session_token", null); + assertEquals(ACCESS_KEY_SETTING, clientSettings.getAllSettings().toArray()[0]); + assertEquals(SECRET_KEY_SETTING, clientSettings.getAllSettings().toArray()[1]); + assertEquals(SESSION_TOKEN_SETTING, clientSettings.getAllSettings().toArray()[2]); + } + + public void testWithSessionCredentials() throws IOException { + PersonalizeClientSettings clientSettings = PersonalizeClientSettingsTestUtil.buildClientSettings(true, true, true); + AWSCredentials credentials = clientSettings.getCredentials(); + assertEquals(ACCESS_KEY, credentials.getAWSAccessKeyId()); + assertEquals(SECRET_KEY, credentials.getAWSSecretKey()); + assertTrue(credentials instanceof AWSSessionCredentials); + AWSSessionCredentials sessionCredentials = (AWSSessionCredentials) credentials; + assertEquals(SESSION_TOKEN, sessionCredentials.getSessionToken()); + } + + public void testWithoutCredentials() throws IOException { + PersonalizeClientSettings clientSettings = PersonalizeClientSettingsTestUtil.buildClientSettings(false, false, false); + assertNull(clientSettings.getCredentials()); + } + + public void testWithoutAccessKey() { + expectThrows(SettingsException.class, () -> PersonalizeClientSettingsTestUtil.buildClientSettings(false, true, false)); + expectThrows(SettingsException.class, () -> PersonalizeClientSettingsTestUtil.buildClientSettings(false, true, true)); + } + + public void testWithoutSecretKey() { + expectThrows(SettingsException.class, () -> PersonalizeClientSettingsTestUtil.buildClientSettings(true, false, false)); + expectThrows(SettingsException.class, () -> PersonalizeClientSettingsTestUtil.buildClientSettings(true, false, true)); + } + + public void testWithSessionTokenButNoCredentials() { + expectThrows(SettingsException.class, () -> PersonalizeClientSettingsTestUtil.buildClientSettings(false, false, true)); + } +} diff --git a/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeClientTests.java b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeClientTests.java new file mode 100644 index 0000000..57c2ead --- /dev/null +++ b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeClientTests.java @@ -0,0 +1,42 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.client; + +import com.amazonaws.auth.AWSCredentials; +import com.amazonaws.auth.AWSCredentialsProvider; +import com.amazonaws.auth.AWSStaticCredentialsProvider; +import com.amazonaws.auth.BasicSessionCredentials; +import com.amazonaws.services.personalizeruntime.model.GetPersonalizedRankingRequest; +import com.amazonaws.services.personalizeruntime.model.GetPersonalizedRankingResult; +import org.mockito.Mockito; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.utils.PersonalizeRuntimeTestUtil; +import org.opensearch.test.OpenSearchTestCase; + +import java.io.IOException; + +import static org.mockito.ArgumentMatchers.any; + +public class PersonalizeClientTests extends OpenSearchTestCase { + + public void testCreateClient() throws IOException { + AWSCredentials credentials = new BasicSessionCredentials("accessKey", "secretKey", "sessionToken"); + AWSCredentialsProvider credentialsProvider = new AWSStaticCredentialsProvider(credentials); + String region = "us-west-2"; + try (PersonalizeClient client = new PersonalizeClient(credentialsProvider,region)) { + assertTrue(client.getPersonalizeRuntime() != null); + } + } + + public void testGetPersonalizedRanking() { + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + GetPersonalizedRankingRequest request = PersonalizeRuntimeTestUtil.buildGetPersonalizedRankingRequest(); + Mockito.when(client.getPersonalizedRanking(any())).thenReturn(PersonalizeRuntimeTestUtil.buildGetPersonalizedRankingResult()); + GetPersonalizedRankingResult result = client.getPersonalizedRanking(request); + assertEquals(result.getRecommendationId(), "sampleRecommendationId"); + } +} diff --git a/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeCredentialsProviderFactoryTests.java b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeCredentialsProviderFactoryTests.java new file mode 100644 index 0000000..af72a46 --- /dev/null +++ b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/client/PersonalizeCredentialsProviderFactoryTests.java @@ -0,0 +1,68 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.client; + +import com.amazonaws.auth.AWSCredentialsProvider; +import com.amazonaws.auth.AWSStaticCredentialsProvider; +import com.amazonaws.auth.DefaultAWSCredentialsProviderChain; +import com.amazonaws.auth.STSAssumeRoleSessionCredentialsProvider; +import com.amazonaws.http.IdleConnectionReaper; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.utils.PersonalizeClientSettingsTestUtil; +import org.opensearch.test.OpenSearchTestCase; + +import java.io.IOException; + +public class PersonalizeCredentialsProviderFactoryTests extends OpenSearchTestCase { + + public void testGetStaticCredentialsProviderWithoutIAMRole() throws IOException { + PersonalizeClientSettings settings = + PersonalizeClientSettingsTestUtil.buildClientSettings(true, true, true); + + AWSCredentialsProvider credentialsProvider = PersonalizeCredentialsProviderFactory.getCredentialsProvider(settings); + assertEquals(credentialsProvider.getClass(), AWSStaticCredentialsProvider.class); + } + + public void testGetDefaultCredentialsProviderWithoutIAMRole() throws IOException { + PersonalizeClientSettings settings = + PersonalizeClientSettingsTestUtil.buildClientSettings(false, false, false); + + AWSCredentialsProvider credentialsProvider = PersonalizeCredentialsProviderFactory.getCredentialsProvider(settings); + assertEquals(credentialsProvider.getClass(), DefaultAWSCredentialsProviderChain.class); + } + + public void testGetCredentialsProviderWithIAMRole() throws IOException { + PersonalizeClientSettings settings = + PersonalizeClientSettingsTestUtil.buildClientSettings(true, true, true); + + String iamRoleArn = "test-iam-role-arn"; + String awsRegion = "us-west-2"; + AWSCredentialsProvider credentialsProvider = PersonalizeCredentialsProviderFactory.getCredentialsProvider(settings, iamRoleArn, awsRegion); + assertEquals(credentialsProvider.getClass(), STSAssumeRoleSessionCredentialsProvider.class); + IdleConnectionReaper.shutdown(); + } + + public void testGetStaticCredentialsProviderWithEmptyIAMRole() throws IOException { + PersonalizeClientSettings settings = + PersonalizeClientSettingsTestUtil.buildClientSettings(true, true, true); + + String iamRoleArn = ""; + String awsRegion = "us-west-2"; + AWSCredentialsProvider credentialsProvider = PersonalizeCredentialsProviderFactory.getCredentialsProvider(settings, iamRoleArn, awsRegion); + assertEquals(credentialsProvider.getClass(), AWSStaticCredentialsProvider.class); + } + + public void testGetDefaultCredentialsProviderWithEmptyIAMRole() throws IOException { + PersonalizeClientSettings settings = + PersonalizeClientSettingsTestUtil.buildClientSettings(false, false, false); + + String iamRoleArn = ""; + String awsRegion = "us-west-2"; + AWSCredentialsProvider credentialsProvider = PersonalizeCredentialsProviderFactory.getCredentialsProvider(settings, iamRoleArn, awsRegion); + assertEquals(credentialsProvider.getClass(), DefaultAWSCredentialsProviderChain.class); + } +} diff --git a/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/configuration/PersonalizeIntelligentRankerConfigurationTests.java b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/configuration/PersonalizeIntelligentRankerConfigurationTests.java new file mode 100644 index 0000000..b54a2e4 --- /dev/null +++ b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/configuration/PersonalizeIntelligentRankerConfigurationTests.java @@ -0,0 +1,32 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration; + +import org.opensearch.test.OpenSearchTestCase; + +public class PersonalizeIntelligentRankerConfigurationTests extends OpenSearchTestCase { + + public void createConfigurationTest() { + String personalizeCampaign = "arn:aws:personalize:us-west-2:000000000000:campaign/test-campaign"; + String iamRoleArn = "sampleRoleArn"; + String recipe = "sample-personalize-recipe"; + String itemIdField = "ITEM_ID"; + String region = "us-west-2"; + double weight = 0.25; + + PersonalizeIntelligentRankerConfiguration config = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, recipe, itemIdField, region, weight); + + assertEquals(config.getPersonalizeCampaign(), personalizeCampaign); + assertEquals(config.getIamRoleArn(), iamRoleArn); + assertEquals(config.getRecipe(), recipe); + assertEquals(config.getItemIdField(), itemIdField); + assertEquals(config.getRegion(), region); + assertEquals(config.getWeight(), weight, 0.0); + } +} diff --git a/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/ranker/PersonalizeRankerFactoryTests.java b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/ranker/PersonalizeRankerFactoryTests.java new file mode 100644 index 0000000..b88f601 --- /dev/null +++ b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/ranker/PersonalizeRankerFactoryTests.java @@ -0,0 +1,47 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.ranker; + +import org.mockito.Mockito; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.client.PersonalizeClient; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration.PersonalizeIntelligentRankerConfiguration; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.reranker.PersonalizedRanker; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.reranker.PersonalizedRankerFactory; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.reranker.impl.AmazonPersonalizedRankerImpl; +import org.opensearch.test.OpenSearchTestCase; + +import static org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration.Constants.AMAZON_PERSONALIZED_RANKING_RECIPE_NAME; + +public class PersonalizeRankerFactoryTests extends OpenSearchTestCase { + + private String personalizeCampaign = "arn:aws:personalize:us-west-2:000000000000:campaign/test-campaign"; + private String iamRoleArn = "sampleRoleArn"; + private String itemIdField = "ITEM_ID"; + private String region = "us-west-2"; + private double weight = 0.25; + + public void testGetPersonalizeRankerForPersonalizedRankingRecipe() { + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, AMAZON_PERSONALIZED_RANKING_RECIPE_NAME, itemIdField, region, weight); + + PersonalizedRankerFactory factory = new PersonalizedRankerFactory(); + PersonalizedRanker ranker = factory.getPersonalizedRanker(rankerConfig, client); + assertEquals(ranker.getClass(), AmazonPersonalizedRankerImpl.class); + } + + public void testGetPersonalizeRankerForUnknownRecipe() { + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, "sample-recipe", itemIdField, region, weight); + + PersonalizedRankerFactory factory = new PersonalizedRankerFactory(); + PersonalizedRanker ranker = factory.getPersonalizedRanker(rankerConfig, client); + assertNull(ranker); + } +} diff --git a/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/ranker/impl/AmazonPersonalizeRankerImplTests.java b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/ranker/impl/AmazonPersonalizeRankerImplTests.java new file mode 100644 index 0000000..1eec5cc --- /dev/null +++ b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/ranker/impl/AmazonPersonalizeRankerImplTests.java @@ -0,0 +1,343 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ + +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.ranker.impl; + +import com.amazonaws.AmazonServiceException; +import org.junit.Assert; +import org.mockito.Mockito; +import org.opensearch.OpenSearchParseException; +import org.opensearch.search.SearchHit; +import org.opensearch.search.SearchHits; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.client.PersonalizeClient; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration.PersonalizeIntelligentRankerConfiguration; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.requestparameter.PersonalizeRequestParameters; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.reranker.impl.AmazonPersonalizedRankerImpl; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.utils.PersonalizeRuntimeTestUtil; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.utils.SearchTestUtil; +import org.opensearch.test.OpenSearchTestCase; + +import java.io.IOException; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.stream.Collectors; + +import static org.mockito.ArgumentMatchers.any; + +public class AmazonPersonalizeRankerImplTests extends OpenSearchTestCase { + + private String personalizeCampaign = "arn:aws:personalize:us-west-2:000000000000:campaign/test-campaign"; + private String iamRoleArn = "sampleRoleArn"; + private String recipe = "sample-personalize-recipe"; + private String itemIdField = "ITEM_ID"; + private String region = "us-west-2"; + private double weight = 0.25; + private int numOfHits = 10; + + public void testReRank() throws IOException { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, recipe, itemIdField, region, weight); + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + Mockito.when(client.getPersonalizedRanking(any())).thenReturn(PersonalizeRuntimeTestUtil.buildGetPersonalizedRankingResult()); + + AmazonPersonalizedRankerImpl ranker = new AmazonPersonalizedRankerImpl(rankerConfig, client); + PersonalizeRequestParameters requestParameters = new PersonalizeRequestParameters(); + requestParameters.setUserId("28"); + SearchHits responseHits = SearchTestUtil.getSampleSearchHitsForPersonalize(numOfHits); + SearchHits transformedHits = ranker.rerank(responseHits, requestParameters); + assertEquals(responseHits.getHits().length, transformedHits.getHits().length); + } + + public void testReRankWithoutItemIdFieldInConfig() throws IOException { + String blankItemIdField = ""; + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, recipe, blankItemIdField, region, weight); + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + Mockito.when(client.getPersonalizedRanking(any())).thenReturn(PersonalizeRuntimeTestUtil.buildGetPersonalizedRankingResult()); + + AmazonPersonalizedRankerImpl ranker = new AmazonPersonalizedRankerImpl(rankerConfig, client); + PersonalizeRequestParameters requestParameters = new PersonalizeRequestParameters(); + requestParameters.setUserId("28"); + SearchHits responseHits = SearchTestUtil.getSampleSearchHitsForPersonalize(numOfHits); + SearchHits transformedHits = ranker.rerank(responseHits, requestParameters); + assertEquals(responseHits.getHits().length, transformedHits.getHits().length); + } + + public void testReRankWithRequestParameterContext() throws IOException { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, recipe, itemIdField, region, weight); + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + Mockito.when(client.getPersonalizedRanking(any())).thenReturn(PersonalizeRuntimeTestUtil.buildGetPersonalizedRankingResult()); + + AmazonPersonalizedRankerImpl ranker = new AmazonPersonalizedRankerImpl(rankerConfig, client); + Map context = new HashMap<>(); + context.put("contextKey", "contextValue"); + PersonalizeRequestParameters requestParameters = new PersonalizeRequestParameters(); + requestParameters.setUserId("28"); + requestParameters.setContext(context); + SearchHits responseHits = SearchTestUtil.getSampleSearchHitsForPersonalize(numOfHits); + SearchHits transformedHits = ranker.rerank(responseHits, requestParameters); + assertEquals(responseHits.getHits().length, transformedHits.getHits().length); + } + + public void testReRankWithInvalidRequestParameterContext() throws IOException { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, recipe, itemIdField, region, weight); + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + Mockito.when(client.getPersonalizedRanking(any())).thenReturn(PersonalizeRuntimeTestUtil.buildGetPersonalizedRankingResult()); + + AmazonPersonalizedRankerImpl ranker = new AmazonPersonalizedRankerImpl(rankerConfig, client); + Map context = new HashMap<>(); + context.put("contextKey", 2); + PersonalizeRequestParameters requestParameters = new PersonalizeRequestParameters(); + requestParameters.setUserId("28"); + requestParameters.setContext(context); + SearchHits responseHits = SearchTestUtil.getSampleSearchHitsForPersonalize(numOfHits); + expectThrows(OpenSearchParseException.class, () -> + ranker.rerank(responseHits, requestParameters)); + } + + public void testReRankWithNoUserId() throws IOException { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, recipe, itemIdField, region, weight); + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + Mockito.when(client.getPersonalizedRanking(any())).thenReturn(PersonalizeRuntimeTestUtil.buildGetPersonalizedRankingResult()); + + AmazonPersonalizedRankerImpl ranker = new AmazonPersonalizedRankerImpl(rankerConfig, client); + Map context = new HashMap<>(); + context.put("contextKey", "contextValue"); + PersonalizeRequestParameters requestParameters = new PersonalizeRequestParameters(); + requestParameters.setContext(context); + SearchHits responseHits = SearchTestUtil.getSampleSearchHitsForPersonalize(numOfHits); + expectThrows(OpenSearchParseException.class, () -> + ranker.rerank(responseHits, requestParameters)); + } + + public void testReRankWithEmptyItemIdField() throws IOException { + String itemIdEmpty = ""; + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, recipe, itemIdEmpty, region, weight); + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + Mockito.when(client.getPersonalizedRanking(any())).thenReturn(PersonalizeRuntimeTestUtil.buildGetPersonalizedRankingResult()); + + AmazonPersonalizedRankerImpl ranker = new AmazonPersonalizedRankerImpl(rankerConfig, client); + PersonalizeRequestParameters requestParameters = new PersonalizeRequestParameters(); + requestParameters.setUserId("28"); + SearchHits responseHits = SearchTestUtil.getSampleSearchHitsForPersonalize(numOfHits); + SearchHits transformedHits = ranker.rerank(responseHits, requestParameters); + assertEquals(responseHits.getHits().length, transformedHits.getHits().length); + } + + public void testReRankWithNullItemIdField() throws IOException { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, recipe, null, region, weight); + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + Mockito.when(client.getPersonalizedRanking(any())).thenReturn(PersonalizeRuntimeTestUtil.buildGetPersonalizedRankingResult()); + + AmazonPersonalizedRankerImpl ranker = new AmazonPersonalizedRankerImpl(rankerConfig, client); + PersonalizeRequestParameters requestParameters = new PersonalizeRequestParameters(); + requestParameters.setUserId("28"); + SearchHits responseHits = SearchTestUtil.getSampleSearchHitsForPersonalize(numOfHits); + SearchHits transformedHits = ranker.rerank(responseHits, requestParameters); + assertEquals(responseHits.getHits().length, transformedHits.getHits().length); + } + + public void testReRankWithWeightAsZero() throws IOException { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, recipe, itemIdField, region, 0); + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + Mockito.when(client.getPersonalizedRanking(any())).thenReturn(PersonalizeRuntimeTestUtil.buildGetPersonalizedRankingResult(numOfHits)); + + AmazonPersonalizedRankerImpl ranker = new AmazonPersonalizedRankerImpl(rankerConfig, client); + PersonalizeRequestParameters requestParameters = new PersonalizeRequestParameters(); + requestParameters.setUserId("28"); + SearchHits responseHits = SearchTestUtil.getSampleSearchHitsForPersonalize(numOfHits); + SearchHits transformedHits = ranker.rerank(responseHits, requestParameters); + + List originalHits = Arrays.asList(transformedHits.getHits()); + String itemIdfield = rankerConfig.getItemIdField(); + List rerankedDocumentIds; + + rerankedDocumentIds = originalHits.stream() + .filter(h -> h.getSourceAsMap().get(itemIdfield) != null) + .map(h -> h.getSourceAsMap().get(itemIdfield).toString()) + .collect(Collectors.toList()); + + ArrayList expectedRankedDocumentIds = PersonalizeRuntimeTestUtil.expectedRankedItemIdsForGivenWeight(numOfHits, 0); + assertEquals(expectedRankedDocumentIds, rerankedDocumentIds); + } + + public void testReRankWithWeightAsOne() throws IOException { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, recipe, itemIdField, region, 1); + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + Mockito.when(client.getPersonalizedRanking(any())).thenReturn(PersonalizeRuntimeTestUtil.buildGetPersonalizedRankingResult(numOfHits)); + + + AmazonPersonalizedRankerImpl ranker = new AmazonPersonalizedRankerImpl(rankerConfig, client); + PersonalizeRequestParameters requestParameters = new PersonalizeRequestParameters(); + requestParameters.setUserId("28"); + SearchHits responseHits = SearchTestUtil.getSampleSearchHitsForPersonalize(numOfHits); + SearchHits transformedHits = ranker.rerank(responseHits, requestParameters); + + List originalHits = Arrays.asList(transformedHits.getHits()); + String itemIdfield = rankerConfig.getItemIdField(); + List rerankedDocumentIds; + + rerankedDocumentIds = originalHits.stream() + .filter(h -> h.getSourceAsMap().get(itemIdfield) != null) + .map(h -> h.getSourceAsMap().get(itemIdfield).toString()) + .collect(Collectors.toList()); + + ArrayList expectedRankedDocumentIds = PersonalizeRuntimeTestUtil.expectedRankedItemIdsForGivenWeight(numOfHits, 1); + assertEquals(expectedRankedDocumentIds, rerankedDocumentIds); + } + + public void testReRankWithWeightAsNeitherZeroOrOne() throws IOException { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, recipe, itemIdField, region, weight); + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + Mockito.when(client.getPersonalizedRanking(any())).thenReturn(PersonalizeRuntimeTestUtil.buildGetPersonalizedRankingResult(numOfHits)); + + AmazonPersonalizedRankerImpl ranker = new AmazonPersonalizedRankerImpl(rankerConfig, client); + PersonalizeRequestParameters requestParameters = new PersonalizeRequestParameters(); + requestParameters.setUserId("28"); + SearchHits responseHits = SearchTestUtil.getSampleSearchHitsForPersonalize(numOfHits); + SearchHits transformedHits = ranker.rerank(responseHits, requestParameters); + + List originalHits = Arrays.asList(transformedHits.getHits()); + String itemIdfield = rankerConfig.getItemIdField(); + List rerankedDocumentIds; + + rerankedDocumentIds = originalHits.stream() + .filter(h -> h.getSourceAsMap().get(itemIdfield) != null) + .map(h -> h.getSourceAsMap().get(itemIdfield).toString()) + .collect(Collectors.toList()); + + ArrayList rerankedDocumentIdsWhenWeightIsOne = PersonalizeRuntimeTestUtil.expectedRankedItemIdsForGivenWeight(numOfHits, 1); + ArrayList rerankedDocumentIdsWhenWeightIsZero = PersonalizeRuntimeTestUtil.expectedRankedItemIdsForGivenWeight(numOfHits, 0); + + assertNotEquals(rerankedDocumentIdsWhenWeightIsOne, rerankedDocumentIds); + assertNotEquals(rerankedDocumentIdsWhenWeightIsZero, rerankedDocumentIds); + } + + public void testReRankWithWeightAsZeroWithNullItemIdField() throws IOException { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, recipe, "", region, 0); + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + Mockito.when(client.getPersonalizedRanking(any())).thenReturn(PersonalizeRuntimeTestUtil.buildGetPersonalizedRankingResult(numOfHits)); + + AmazonPersonalizedRankerImpl ranker = new AmazonPersonalizedRankerImpl(rankerConfig, client); + PersonalizeRequestParameters requestParameters = new PersonalizeRequestParameters(); + requestParameters.setUserId("28"); + SearchHits responseHits = SearchTestUtil.getSampleSearchHitsForPersonalize(numOfHits); + SearchHits transformedHits = ranker.rerank(responseHits, requestParameters); + + List originalHits = Arrays.asList(transformedHits.getHits()); + List rerankedDocumentIds; + + rerankedDocumentIds = originalHits.stream() + .filter(h -> h.getId() != null) + .map(h -> h.getId()) + .collect(Collectors.toList()); + + ArrayList expectedRankedDocumentIds = PersonalizeRuntimeTestUtil.expectedRankedItemIdsForGivenWeight(numOfHits, 0); + assertEquals(expectedRankedDocumentIds, rerankedDocumentIds); + } + + + public void testReRankWithWeightAsOneWithNullItemIdField() throws IOException { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, recipe, "", region, 1); + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + Mockito.when(client.getPersonalizedRanking(any())).thenReturn(PersonalizeRuntimeTestUtil.buildGetPersonalizedRankingResult(numOfHits)); + + + AmazonPersonalizedRankerImpl ranker = new AmazonPersonalizedRankerImpl(rankerConfig, client); + PersonalizeRequestParameters requestParameters = new PersonalizeRequestParameters(); + requestParameters.setUserId("28"); + SearchHits responseHits = SearchTestUtil.getSampleSearchHitsForPersonalize(numOfHits); + SearchHits transformedHits = ranker.rerank(responseHits, requestParameters); + + List originalHits = Arrays.asList(transformedHits.getHits()); + List rerankedDocumentIds; + + rerankedDocumentIds = originalHits.stream() + .filter(h -> h.getId() != null) + .map(h -> h.getId()) + .collect(Collectors.toList()); + + ArrayList expectedRankedDocumentIds = PersonalizeRuntimeTestUtil.expectedRankedItemIdsForGivenWeight(numOfHits, 1); + assertEquals(expectedRankedDocumentIds, rerankedDocumentIds); + } + + public void testReRankWithWeightAsNeitherZeroOrOneWithNullItemIdField() throws IOException { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, recipe, "", region, weight); + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + Mockito.when(client.getPersonalizedRanking(any())).thenReturn(PersonalizeRuntimeTestUtil.buildGetPersonalizedRankingResult(numOfHits)); + + AmazonPersonalizedRankerImpl ranker = new AmazonPersonalizedRankerImpl(rankerConfig, client); + PersonalizeRequestParameters requestParameters = new PersonalizeRequestParameters(); + requestParameters.setUserId("28"); + SearchHits responseHits = SearchTestUtil.getSampleSearchHitsForPersonalize(numOfHits); + SearchHits transformedHits = ranker.rerank(responseHits, requestParameters); + + List originalHits = Arrays.asList(transformedHits.getHits()); + List rerankedDocumentIds; + + rerankedDocumentIds = originalHits.stream() + .filter(h -> h.getId() != null) + .map(h -> h.getId()) + .collect(Collectors.toList()); + + ArrayList rerankedDocumentIdsWhenWeightIsOne = PersonalizeRuntimeTestUtil.expectedRankedItemIdsForGivenWeight(numOfHits, 1); + ArrayList rerankedDocumentIdsWhenWeightIsZero = PersonalizeRuntimeTestUtil.expectedRankedItemIdsForGivenWeight(numOfHits, 0); + + assertNotEquals(rerankedDocumentIdsWhenWeightIsOne, rerankedDocumentIds); + assertNotEquals(rerankedDocumentIdsWhenWeightIsZero, rerankedDocumentIds); + } + + public void testReRankWithaccessDeniedExceptionWithStatusCode400() throws IOException { + + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, recipe, itemIdField, region, weight); + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + Mockito.when(client.getPersonalizedRanking(any())).thenThrow(buildErrorWithStatusCode(400)); + + PersonalizeRequestParameters requestParameters = new PersonalizeRequestParameters(); + requestParameters.setUserId("28"); + SearchHits responseHits = SearchTestUtil.getSampleSearchHitsForPersonalize(numOfHits); + AmazonPersonalizedRankerImpl ranker = new AmazonPersonalizedRankerImpl(rankerConfig, client); + Assert.assertThrows(IllegalArgumentException.class, () -> ranker.rerank(responseHits, requestParameters)); + } + + public void testReRankWithaccessDeniedExceptionWithStatusCode500() throws IOException { + + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, recipe, itemIdField, region, weight); + PersonalizeClient client = Mockito.mock(PersonalizeClient.class); + Mockito.when(client.getPersonalizedRanking(any())).thenThrow(buildErrorWithStatusCode(500)); + + PersonalizeRequestParameters requestParameters = new PersonalizeRequestParameters(); + requestParameters.setUserId("28"); + SearchHits responseHits = SearchTestUtil.getSampleSearchHitsForPersonalize(numOfHits); + AmazonPersonalizedRankerImpl ranker = new AmazonPersonalizedRankerImpl(rankerConfig, client); + Assert.assertThrows(AmazonServiceException.class, () -> ranker.rerank(responseHits, requestParameters)); + } + + + private AmazonServiceException buildErrorWithStatusCode(int statusCode) { + AmazonServiceException amazonServiceException = new AmazonServiceException("Error"); + amazonServiceException.setStatusCode(statusCode); + return amazonServiceException; + } +} diff --git a/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/requestparameter/PersonalizeRequestParameterUtilTests.java b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/requestparameter/PersonalizeRequestParameterUtilTests.java new file mode 100644 index 0000000..ea09d79 --- /dev/null +++ b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/requestparameter/PersonalizeRequestParameterUtilTests.java @@ -0,0 +1,97 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + * + */ + +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.requestparameter; + +import org.opensearch.action.search.SearchRequest; +import org.opensearch.search.builder.SearchSourceBuilder; +import org.opensearch.test.OpenSearchTestCase; + +import java.util.HashMap; +import java.util.List; +import java.util.Map; + +public class PersonalizeRequestParameterUtilTests extends OpenSearchTestCase { + + public void testExtractParameters() { + PersonalizeRequestParameters expected = new PersonalizeRequestParameters("user_1", new HashMap<>()); + PersonalizeRequestParametersExtBuilder extBuilder = new PersonalizeRequestParametersExtBuilder(); + extBuilder.setRequestParameters(expected); + SearchSourceBuilder sourceBuilder = SearchSourceBuilder.searchSource() + .ext(List.of(extBuilder)); + SearchRequest request = new SearchRequest("my_index").source(sourceBuilder); + PersonalizeRequestParameters actual = PersonalizeRequestParameterUtil.getPersonalizeRequestParameters(request); + assertEquals(expected, actual); + } + + public void testExtractParametersWithContext() { + Map context = new HashMap<>(); + context.put("contextKey", "contextValue"); + PersonalizeRequestParameters expected = new PersonalizeRequestParameters("user_1", context); + PersonalizeRequestParametersExtBuilder extBuilder = new PersonalizeRequestParametersExtBuilder(); + extBuilder.setRequestParameters(expected); + SearchSourceBuilder sourceBuilder = SearchSourceBuilder.searchSource() + .ext(List.of(extBuilder)); + SearchRequest request = new SearchRequest("my_index").source(sourceBuilder); + PersonalizeRequestParameters actual = PersonalizeRequestParameterUtil.getPersonalizeRequestParameters(request); + assertEquals(expected, actual); + } + + public void testPersonalizeRequestParametersEquals() { + Map notExpectedContext = new HashMap<>(); + notExpectedContext.put("contextKey", "contextValue"); + PersonalizeRequestParameters notExpected = new PersonalizeRequestParameters("user_1", notExpectedContext); + + Map expectedContext = new HashMap<>(); + expectedContext.put("contextKey2", "contextValue2"); + PersonalizeRequestParameters expected = new PersonalizeRequestParameters("user_1", expectedContext); + PersonalizeRequestParametersExtBuilder extBuilder = new PersonalizeRequestParametersExtBuilder(); + extBuilder.setRequestParameters(expected); + SearchSourceBuilder sourceBuilder = SearchSourceBuilder.searchSource() + .ext(List.of(extBuilder)); + SearchRequest request = new SearchRequest("my_index").source(sourceBuilder); + PersonalizeRequestParameters actual = PersonalizeRequestParameterUtil.getPersonalizeRequestParameters(request); + assertNotEquals(notExpected, actual); + } + + public void testPersonalizeRequestParametersContextMapDifferentSize() { + Map notExpectedContext = new HashMap<>(); + notExpectedContext.put("contextKey", "contextValue"); + PersonalizeRequestParameters notExpected = new PersonalizeRequestParameters("user_1", notExpectedContext); + + Map expectedContext = new HashMap<>(); + expectedContext.put("contextKey2", "contextValue2"); + expectedContext.put("contextKey22", "contextValue22"); + PersonalizeRequestParameters expected = new PersonalizeRequestParameters("user_1", expectedContext); + PersonalizeRequestParametersExtBuilder extBuilder = new PersonalizeRequestParametersExtBuilder(); + extBuilder.setRequestParameters(expected); + SearchSourceBuilder sourceBuilder = SearchSourceBuilder.searchSource() + .ext(List.of(extBuilder)); + SearchRequest request = new SearchRequest("my_index").source(sourceBuilder); + PersonalizeRequestParameters actual = PersonalizeRequestParameterUtil.getPersonalizeRequestParameters(request); + assertNotEquals(notExpected, actual); + } + + public void testPersonalizeRequestParametersUserIdDiffers() { + Map notExpectedContext = new HashMap<>(); + notExpectedContext.put("contextKey", "contextValue"); + PersonalizeRequestParameters notExpected = new PersonalizeRequestParameters("user_1", notExpectedContext); + + Map expectedContext = new HashMap<>(); + expectedContext.put("contextKey", "contextValue"); + PersonalizeRequestParameters expected = new PersonalizeRequestParameters("user_2", expectedContext); + PersonalizeRequestParametersExtBuilder extBuilder = new PersonalizeRequestParametersExtBuilder(); + extBuilder.setRequestParameters(expected); + SearchSourceBuilder sourceBuilder = SearchSourceBuilder.searchSource() + .ext(List.of(extBuilder)); + SearchRequest request = new SearchRequest("my_index").source(sourceBuilder); + PersonalizeRequestParameters actual = PersonalizeRequestParameterUtil.getPersonalizeRequestParameters(request); + assertNotEquals(notExpected, actual); + } +} \ No newline at end of file diff --git a/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/requestparameter/PersonalizeRequestParametersExtBuilderTests.java b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/requestparameter/PersonalizeRequestParametersExtBuilderTests.java new file mode 100644 index 0000000..fc1c8c9 --- /dev/null +++ b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/requestparameter/PersonalizeRequestParametersExtBuilderTests.java @@ -0,0 +1,55 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.requestparameter; + +import org.opensearch.core.common.bytes.BytesReference; +import org.opensearch.common.io.stream.BytesStreamOutput; +import org.opensearch.common.xcontent.XContentHelper; +import org.opensearch.common.xcontent.XContentType; +import org.opensearch.core.xcontent.XContentParser; +import org.opensearch.test.OpenSearchTestCase; + +import java.io.IOException; +import java.util.HashMap; +import java.util.Map; + +public class PersonalizeRequestParametersExtBuilderTests extends OpenSearchTestCase { + + public void testXContentRoundTrip() throws IOException { + Map context = new HashMap<>(); + context.put("contextKey", "contextValue"); + PersonalizeRequestParameters requestParameters = new PersonalizeRequestParameters("28", context); + PersonalizeRequestParametersExtBuilder personalizeExtBuilder = new PersonalizeRequestParametersExtBuilder(); + personalizeExtBuilder.setRequestParameters(requestParameters); + XContentType xContentType = randomFrom(XContentType.values()); + BytesReference serialized = XContentHelper.toXContent(personalizeExtBuilder, xContentType, true); + + XContentParser parser = createParser(xContentType.xContent(), serialized); + + PersonalizeRequestParametersExtBuilder deserialized = + PersonalizeRequestParametersExtBuilder.parse(parser); + + assertEquals(personalizeExtBuilder, deserialized); + } + + public void testStreamRoundTrip() throws IOException { + PersonalizeRequestParameters requestParameters = new PersonalizeRequestParameters(); + requestParameters.setUserId("28"); + requestParameters.setContext(new HashMap<>()); + PersonalizeRequestParametersExtBuilder personalizeExtBuilder = new PersonalizeRequestParametersExtBuilder(); + personalizeExtBuilder.setRequestParameters(requestParameters); + BytesStreamOutput bytesStreamOutput = new BytesStreamOutput(); + personalizeExtBuilder.writeTo(bytesStreamOutput); + + PersonalizeRequestParametersExtBuilder deserialized = + new PersonalizeRequestParametersExtBuilder(bytesStreamOutput.bytes().streamInput()); + assertEquals(personalizeExtBuilder, deserialized); + } + + +} diff --git a/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/utils/PersonalizeClientSettingsTestUtil.java b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/utils/PersonalizeClientSettingsTestUtil.java new file mode 100644 index 0000000..d5e4052 --- /dev/null +++ b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/utils/PersonalizeClientSettingsTestUtil.java @@ -0,0 +1,39 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.utils; + +import org.opensearch.common.settings.MockSecureSettings; +import org.opensearch.common.settings.Settings; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.client.PersonalizeClientSettings; + +import java.io.IOException; + +public class PersonalizeClientSettingsTestUtil { + public static final String ACCESS_KEY = "my-access-key"; + public static final String SECRET_KEY = "my-secret-key"; + public static final String SESSION_TOKEN = "session-token"; + + public static PersonalizeClientSettings buildClientSettings(boolean withAccessKey, boolean withSecretKey, + boolean withSessionToken) throws IOException { + try (MockSecureSettings secureSettings = new MockSecureSettings()) { + if (withAccessKey) { + secureSettings.setString(PersonalizeClientSettings.ACCESS_KEY_SETTING.getKey(), ACCESS_KEY); + } + if (withSecretKey) { + secureSettings.setString(PersonalizeClientSettings.SECRET_KEY_SETTING.getKey(), SECRET_KEY); + } + if (withSessionToken) { + secureSettings.setString(PersonalizeClientSettings.SESSION_TOKEN_SETTING.getKey(), SESSION_TOKEN); + } + Settings settings = Settings.builder() + .setSecureSettings(secureSettings) + .build(); + return PersonalizeClientSettings.getClientSettings(settings); + } + } +} diff --git a/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/utils/PersonalizeRuntimeTestUtil.java b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/utils/PersonalizeRuntimeTestUtil.java new file mode 100644 index 0000000..4028c3d --- /dev/null +++ b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/utils/PersonalizeRuntimeTestUtil.java @@ -0,0 +1,81 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.utils; + +import com.amazonaws.services.personalizeruntime.model.GetPersonalizedRankingRequest; +import com.amazonaws.services.personalizeruntime.model.GetPersonalizedRankingResult; +import com.amazonaws.services.personalizeruntime.model.PredictedItem; +import org.mockito.Mockito; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.client.PersonalizeClient; + +import java.util.ArrayList; +import java.util.List; +import java.util.function.Function; + +import static org.mockito.ArgumentMatchers.any; + +public class PersonalizeRuntimeTestUtil { + + public static GetPersonalizedRankingRequest buildGetPersonalizedRankingRequest() { + GetPersonalizedRankingRequest request = new GetPersonalizedRankingRequest() + .withUserId("sampleUserId") + .withInputList(new ArrayList()) + .withCampaignArn("sampleCampaign"); + return request; + } + + public static GetPersonalizedRankingResult buildGetPersonalizedRankingResult() { + List predictedItems = new ArrayList<>(); + GetPersonalizedRankingResult result = new GetPersonalizedRankingResult() + .withPersonalizedRanking(predictedItems) + .withRecommendationId("sampleRecommendationId"); + return result; + } + + public static GetPersonalizedRankingResult buildGetPersonalizedRankingResult(int numOfHits) { + List predictedItems = new ArrayList<>(); + for(int i = numOfHits; i >= 1; i--){ + PredictedItem predictedItem = new PredictedItem(). + withScore((double) i/10). + withItemId(String.valueOf(i-1)); + predictedItems.add(predictedItem); + } + GetPersonalizedRankingResult result = new GetPersonalizedRankingResult() + .withPersonalizedRanking(predictedItems) + .withRecommendationId("sampleRecommendationId"); + return result; + } + + public static ArrayList expectedRankedItemIdsForGivenWeight(int numOfHits, int weight){ + ArrayList expectedRankedItemIds = new ArrayList<>(); + if (weight == 0) { + for (int i = 0; i < numOfHits; i++) { + expectedRankedItemIds.add(String.valueOf(i)); + } + } else if (weight == 1){ + for(int i = numOfHits; i >= 1; i--){ + expectedRankedItemIds.add(String.valueOf(i-1)); + } + } + return expectedRankedItemIds; + } + + public static PersonalizeClient buildMockPersonalizeClient() { + return buildMockPersonalizeClient(r -> buildGetPersonalizedRankingResult(10)); + } + + private static PersonalizeClient buildMockPersonalizeClient( + Function mockGetPersonalizedRankingImpl) { + PersonalizeClient personalizeClient = Mockito.mock(PersonalizeClient.class); + Mockito.doAnswer(invocation -> { + GetPersonalizedRankingRequest request = invocation.getArgument(0); + return mockGetPersonalizedRankingImpl.apply(request); + }).when(personalizeClient).getPersonalizedRanking(any(GetPersonalizedRankingRequest.class)); + return personalizeClient; + } +} diff --git a/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/utils/SearchTestUtil.java b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/utils/SearchTestUtil.java new file mode 100644 index 0000000..ca08126 --- /dev/null +++ b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/utils/SearchTestUtil.java @@ -0,0 +1,55 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.utils; + +import org.apache.lucene.search.TotalHits; +import org.opensearch.action.search.SearchRequest; +import org.opensearch.core.common.bytes.BytesReference; +import org.opensearch.common.xcontent.json.JsonXContent; +import org.opensearch.core.xcontent.XContentBuilder; +import org.opensearch.search.SearchHit; +import org.opensearch.search.SearchHits; +import org.opensearch.search.builder.SearchSourceBuilder; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.requestparameter.PersonalizeRequestParameters; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.requestparameter.PersonalizeRequestParametersExtBuilder; + +import java.io.IOException; +import java.util.List; +import java.util.Map; + +public class SearchTestUtil { + public static SearchHits getSampleSearchHitsForPersonalize(int numHits) throws IOException { + SearchHit[] hitsArray = new SearchHit[numHits]; + float maxScore = 0.0f; + for (int i = 0; i < numHits; i++) { + XContentBuilder sourceContent = JsonXContent.contentBuilder() + .startObject() + .field("ITEM_ID", String.valueOf(i)) + .field("body", "Body text for document number " + i) + .field("title", "This is the title for document " + i) + .endObject(); + hitsArray[i] = new SearchHit(i, String.valueOf(i), Map.of(), Map.of()); + float score = (float)(numHits-i)/10; + maxScore = Math.max(score, maxScore); + hitsArray[i].score(score); + hitsArray[i].sourceRef(BytesReference.bytes(sourceContent)); + } + SearchHits searchHits = new SearchHits(hitsArray, new TotalHits(numHits, TotalHits.Relation.EQUAL_TO), maxScore); + return searchHits; + } + + public static SearchRequest createSearchRequestWithPersonalizeRequest(PersonalizeRequestParameters personalizeRequestParams) { + PersonalizeRequestParametersExtBuilder extBuilder = new PersonalizeRequestParametersExtBuilder(); + extBuilder.setRequestParameters(personalizeRequestParams); + SearchSourceBuilder sourceBuilder = SearchSourceBuilder.searchSource() + .ext(List.of(extBuilder)); + + SearchRequest searchRequest = new SearchRequest().source(sourceBuilder); + return searchRequest; + } +} diff --git a/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/utils/ValidationUtilTests.java b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/utils/ValidationUtilTests.java new file mode 100644 index 0000000..1227915 --- /dev/null +++ b/amazon-personalize-ranking/src/test/java/org/opensearch/search/relevance/transformer/personalizeintelligentranking/utils/ValidationUtilTests.java @@ -0,0 +1,101 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ +package org.opensearch.search.relevance.transformer.personalizeintelligentranking.utils; + +import com.amazonaws.http.IdleConnectionReaper; +import org.opensearch.OpenSearchParseException; +import org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration.PersonalizeIntelligentRankerConfiguration; +import org.opensearch.test.OpenSearchTestCase; + +import static org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration.Constants.AMAZON_PERSONALIZED_RANKING_RECIPE_NAME; +import static org.opensearch.search.relevance.transformer.personalizeintelligentranking.configuration.Constants.AMAZON_PERSONALIZED_RANKING_V2_RECIPE_NAME; + +public class ValidationUtilTests extends OpenSearchTestCase { + + private static final String TYPE = "personalize_ranking"; + private static final String TAG = "test_tag"; + private String personalizeCampaign = "arn:aws:personalize:us-west-2:000000000000:campaign/test-campaign"; + private String iamRoleArn = "arn:aws:iam::000000000000:role/test"; + private String itemIdField = "ITEM_ID"; + private String region = "us-west-2"; + private double weight = 1.0; + + public void testValidRankerConfig () { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, AMAZON_PERSONALIZED_RANKING_RECIPE_NAME, itemIdField, region, weight); + ValidationUtil.validatePersonalizeIntelligentRankerConfiguration(rankerConfig, TYPE, TAG); + } + + public void testValidRankerConfigPersonalizedRankingV2 () { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, AMAZON_PERSONALIZED_RANKING_V2_RECIPE_NAME, itemIdField, region, weight); + ValidationUtil.validatePersonalizeIntelligentRankerConfiguration(rankerConfig, TYPE, TAG); + } + + public void testInvalidCampaignArn () { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration("invalid:campaign/test", iamRoleArn, AMAZON_PERSONALIZED_RANKING_RECIPE_NAME, itemIdField, region, weight); + expectThrows(OpenSearchParseException.class, () -> + ValidationUtil.validatePersonalizeIntelligentRankerConfiguration(rankerConfig, TYPE, TAG)); + } + + public void testEmptyCampaignArn () { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration("", iamRoleArn, AMAZON_PERSONALIZED_RANKING_RECIPE_NAME, itemIdField, region, weight); + expectThrows(OpenSearchParseException.class, () -> + ValidationUtil.validatePersonalizeIntelligentRankerConfiguration(rankerConfig, TYPE, TAG)); + } + + public void testNonPersonalizeArnAsCampaignArn () { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration("arn:aws:es:us-west-2:000000000000:domain/testmovies", iamRoleArn, AMAZON_PERSONALIZED_RANKING_RECIPE_NAME, itemIdField, region, weight); + expectThrows(OpenSearchParseException.class, () -> + ValidationUtil.validatePersonalizeIntelligentRankerConfiguration(rankerConfig, TYPE, TAG)); + } + + public void testInvalidIamRoleArn () { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, "invalid:arn/test", AMAZON_PERSONALIZED_RANKING_RECIPE_NAME, itemIdField, region, weight); + expectThrows(OpenSearchParseException.class, () -> + ValidationUtil.validatePersonalizeIntelligentRankerConfiguration(rankerConfig, TYPE, TAG)); + } + + public void testNonIamArnAsIamRoleArn () { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, "arn:aws:es:us-west-2:000000000000:domain/testmovies", AMAZON_PERSONALIZED_RANKING_RECIPE_NAME, itemIdField, region, weight); + expectThrows(OpenSearchParseException.class, () -> + ValidationUtil.validatePersonalizeIntelligentRankerConfiguration(rankerConfig, TYPE, TAG)); + } + + public void testEmptyIamRoleArnAllowed () { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, "", AMAZON_PERSONALIZED_RANKING_RECIPE_NAME, itemIdField, region, weight); + ValidationUtil.validatePersonalizeIntelligentRankerConfiguration(rankerConfig, TYPE, TAG); + } + + public void testInvalidWeightValueGreaterThanRange () { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, AMAZON_PERSONALIZED_RANKING_RECIPE_NAME, itemIdField, region, 3.0); + expectThrows(OpenSearchParseException.class, () -> + ValidationUtil.validatePersonalizeIntelligentRankerConfiguration(rankerConfig, TYPE, TAG)); + } + + public void testInvalidWeightValueLessThanRange () { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, AMAZON_PERSONALIZED_RANKING_RECIPE_NAME, itemIdField, region, -1.0); + expectThrows(OpenSearchParseException.class, () -> + ValidationUtil.validatePersonalizeIntelligentRankerConfiguration(rankerConfig, TYPE, TAG)); + } + + public void testNonPersonalizedRankingRecipeConfig () { + PersonalizeIntelligentRankerConfiguration rankerConfig = + new PersonalizeIntelligentRankerConfiguration(personalizeCampaign, iamRoleArn, "aws-user-personalization", itemIdField, region, -1.0); + expectThrows(OpenSearchParseException.class, () -> + ValidationUtil.validatePersonalizeIntelligentRankerConfiguration(rankerConfig, TYPE, TAG)); + } +} diff --git a/src/yamlRestTest/java/org/opensearch/search/relevance/SearchRelevanceClientYamlTestSuiteIT.java b/amazon-personalize-ranking/src/yamlRestTest/java/org/opensearch/search/relevance/AmazonPersonalizeRankingClientYamlTestSuiteIT.java similarity index 77% rename from src/yamlRestTest/java/org/opensearch/search/relevance/SearchRelevanceClientYamlTestSuiteIT.java rename to amazon-personalize-ranking/src/yamlRestTest/java/org/opensearch/search/relevance/AmazonPersonalizeRankingClientYamlTestSuiteIT.java index 1f0d7dd..43e0349 100644 --- a/src/yamlRestTest/java/org/opensearch/search/relevance/SearchRelevanceClientYamlTestSuiteIT.java +++ b/amazon-personalize-ranking/src/yamlRestTest/java/org/opensearch/search/relevance/AmazonPersonalizeRankingClientYamlTestSuiteIT.java @@ -13,9 +13,9 @@ import org.opensearch.test.rest.yaml.OpenSearchClientYamlSuiteTestCase; -public class SearchRelevanceClientYamlTestSuiteIT extends OpenSearchClientYamlSuiteTestCase { +public class AmazonPersonalizeRankingClientYamlTestSuiteIT extends OpenSearchClientYamlSuiteTestCase { - public SearchRelevanceClientYamlTestSuiteIT(@Name("yaml") ClientYamlTestCandidate testCandidate) { + public AmazonPersonalizeRankingClientYamlTestSuiteIT(@Name("yaml") ClientYamlTestCandidate testCandidate) { super(testCandidate); } diff --git a/amazon-personalize-ranking/src/yamlRestTest/resources/rest-api-spec/test/10_basic.yml b/amazon-personalize-ranking/src/yamlRestTest/resources/rest-api-spec/test/10_basic.yml new file mode 100644 index 0000000..daa9333 --- /dev/null +++ b/amazon-personalize-ranking/src/yamlRestTest/resources/rest-api-spec/test/10_basic.yml @@ -0,0 +1,17 @@ +"Test that the plugin is loaded in OpenSearch": + - do: + cat.plugins: + local: true + h: component + + - match: + $body: /^opensearch-amazon-personalize-ranking-\d+.\d+.\d+.\d+\n$/ + + - do: + indices.create: + index: test + + - do: + search: + index: test + body: { } diff --git a/build.gradle b/build.gradle index 934b9e2..3647cd1 100644 --- a/build.gradle +++ b/build.gradle @@ -1,147 +1,8 @@ -import org.opensearch.gradle.test.RestIntegTestTask - -apply plugin: 'java' -apply plugin: 'idea' -apply plugin: 'opensearch.opensearchplugin' -apply plugin: 'opensearch.yaml-rest-test' -apply plugin: 'opensearch.pluginzip' -apply plugin: 'jacoco' - -group = 'org.opensearch' - -def pluginName = 'search-processor' -def pluginDescription = 'Make Opensearch results more relevant.' -def projectPath = 'org.opensearch' -def pathToPlugin = 'search.relevance' -def pluginClassName = 'SearchRelevancePlugin' - -publishing { - publications { - pluginZip(MavenPublication) { publication -> - pom { - name = pluginName - description = pluginDescription - licenses { - license { - name = "The Apache License, Version 2.0" - url = "http://www.apache.org/licenses/LICENSE-2.0.txt" - } - } - developers { - developer { - name = "OpenSearch" - url = "https://github.com/opensearch-project/search-processor" - } - } - } - } +ext { + isSnapshot = "true" == System.getProperty("build.snapshot", "true") + opensearch_version = System.getProperty("opensearch.version", "3.0.0") + plugin_version = opensearch_version + if (isSnapshot) { + opensearch_version += "-SNAPSHOT" } } -opensearchplugin { - name "opensearch-${pluginName}-${plugin_version}.0" - description pluginDescription - classname "${projectPath}.${pathToPlugin}.${pluginClassName}" - licenseFile rootProject.file('LICENSE') - noticeFile rootProject.file('NOTICE') -} - -// This requires an additional Jar not published as part of build-tools -loggerUsageCheck.enabled = false - -// No need to validate pom, as we do not upload to maven/sonatype -validateNebulaPom.enabled = false - -buildscript { - ext { - isSnapshot = "true" == System.getProperty("build.snapshot", "true") - opensearch_version = System.getProperty("opensearch.version", "3.0.0") - plugin_version = opensearch_version - if (isSnapshot) { - opensearch_version += "-SNAPSHOT" - } - } - - repositories { - mavenLocal() - maven { url "https://aws.oss.sonatype.org/content/repositories/snapshots" } - mavenCentral() - maven { url "https://plugins.gradle.org/m2/" } - } - - dependencies { - classpath "org.opensearch.gradle:build-tools:${opensearch_version}" - } -} - -repositories { - mavenLocal() - maven { url "https://aws.oss.sonatype.org/content/repositories/snapshots" } - mavenCentral() - maven { url "https://plugins.gradle.org/m2/" } -} - -dependencies { - implementation 'com.ibm.icu:icu4j:57.2' - implementation 'org.apache.commons:commons-lang3:3.12.0' - implementation 'org.apache.httpcomponents:httpclient:4.5.13' - implementation 'org.apache.httpcomponents:httpcore:4.4.15' - implementation 'com.fasterxml.jackson.core:jackson-databind:2.14.1' - implementation 'com.fasterxml.jackson.core:jackson-core:2.14.2' - implementation 'com.fasterxml.jackson.core:jackson-annotations:2.14.1' - implementation 'commons-logging:commons-logging:1.2' - implementation 'com.amazonaws:aws-java-sdk-sts:1.12.300' - implementation 'com.amazonaws:aws-java-sdk-core:1.12.300' -} - -test { - include '**/*Tests.class' - finalizedBy jacocoTestReport -} - -task integTest(type: RestIntegTestTask) { - description = "Run tests against a cluster" - testClassesDirs = sourceSets.test.output.classesDirs - classpath = sourceSets.test.runtimeClasspath -} -tasks.named("check").configure { dependsOn(integTest) } - -integTest { - // The --debug-jvm command-line option makes the cluster debuggable; this makes the tests debuggable - if (System.getProperty("test.debug") != null) { - jvmArgs '-agentlib:jdwp=transport=dt_socket,server=y,suspend=y,address=*:5005' - } -} - -testClusters.integTest { - testDistribution = "INTEG_TEST" - - // This installs our plugin into the testClusters - plugin(project.tasks.bundlePlugin.archiveFile) -} - -run { - useCluster testClusters.integTest -} - -sourceSets { - main { - resources { - srcDirs = ["config"] - includes = ["**/*.yml"] - } - } -} - - -jacocoTestReport { - reports { - xml.enabled = true - html.enabled = true - } - dependsOn test -} - -// TODO: Enable these checks -dependencyLicenses.enabled = false -thirdPartyAudit.enabled = false -loggerUsageCheck.enabled = false diff --git a/gradle/wrapper/gradle-wrapper.jar b/gradle/wrapper/gradle-wrapper.jar index 41d9927..a4b76b9 100644 Binary files a/gradle/wrapper/gradle-wrapper.jar and b/gradle/wrapper/gradle-wrapper.jar differ diff --git a/gradle/wrapper/gradle-wrapper.properties b/gradle/wrapper/gradle-wrapper.properties index aa991fc..e1b837a 100644 --- a/gradle/wrapper/gradle-wrapper.properties +++ b/gradle/wrapper/gradle-wrapper.properties @@ -1,5 +1,8 @@ distributionBase=GRADLE_USER_HOME distributionPath=wrapper/dists -distributionUrl=https\://services.gradle.org/distributions/gradle-7.4.2-bin.zip +distributionSha256Sum=7a00d51fb93147819aab76024feece20b6b84e420694101f276be952e08bef03 +distributionUrl=https\://services.gradle.org/distributions/gradle-8.12-bin.zip +networkTimeout=10000 +validateDistributionUrl=true zipStoreBase=GRADLE_USER_HOME zipStorePath=wrapper/dists diff --git a/gradlew b/gradlew index 005bcde..f5feea6 100755 --- a/gradlew +++ b/gradlew @@ -15,6 +15,8 @@ # See the License for the specific language governing permissions and # limitations under the License. # +# SPDX-License-Identifier: Apache-2.0 +# ############################################################################## # @@ -55,7 +57,7 @@ # Darwin, MinGW, and NonStop. # # (3) This script is generated from the Groovy template -# https://github.com/gradle/gradle/blob/master/subprojects/plugins/src/main/resources/org/gradle/api/internal/plugins/unixStartScript.txt +# https://github.com/gradle/gradle/blob/HEAD/platforms/jvm/plugins-application/src/main/resources/org/gradle/api/internal/plugins/unixStartScript.txt # within the Gradle project. # # You can find Gradle at https://github.com/gradle/gradle/. @@ -80,13 +82,12 @@ do esac done -APP_HOME=$( cd "${APP_HOME:-./}" && pwd -P ) || exit - -APP_NAME="Gradle" +# This is normally unused +# shellcheck disable=SC2034 APP_BASE_NAME=${0##*/} - -# Add default JVM options here. You can also use JAVA_OPTS and GRADLE_OPTS to pass JVM options to this script. -DEFAULT_JVM_OPTS='-Dfile.encoding=UTF-8 "-Xmx64m" "-Xms64m"' +# Discard cd standard output in case $CDPATH is set (https://github.com/gradle/gradle/issues/25036) +APP_HOME=$( cd -P "${APP_HOME:-./}" > /dev/null && printf '%s +' "$PWD" ) || exit # Use the maximum available, or set MAX_FD != -1 to use that value. MAX_FD=maximum @@ -133,22 +134,29 @@ location of your Java installation." fi else JAVACMD=java - which java >/dev/null 2>&1 || die "ERROR: JAVA_HOME is not set and no 'java' command could be found in your PATH. + if ! command -v java >/dev/null 2>&1 + then + die "ERROR: JAVA_HOME is not set and no 'java' command could be found in your PATH. Please set the JAVA_HOME variable in your environment to match the location of your Java installation." + fi fi # Increase the maximum file descriptors if we can. if ! "$cygwin" && ! "$darwin" && ! "$nonstop" ; then case $MAX_FD in #( max*) + # In POSIX sh, ulimit -H is undefined. That's why the result is checked to see if it worked. + # shellcheck disable=SC2039,SC3045 MAX_FD=$( ulimit -H -n ) || warn "Could not query maximum file descriptor limit" esac case $MAX_FD in #( '' | soft) :;; #( *) + # In POSIX sh, ulimit -n is undefined. That's why the result is checked to see if it worked. + # shellcheck disable=SC2039,SC3045 ulimit -n "$MAX_FD" || warn "Could not set maximum file descriptor limit to $MAX_FD" esac @@ -193,11 +201,15 @@ if "$cygwin" || "$msys" ; then done fi -# Collect all arguments for the java command; -# * $DEFAULT_JVM_OPTS, $JAVA_OPTS, and $GRADLE_OPTS can contain fragments of -# shell script including quotes and variable substitutions, so put them in -# double quotes to make sure that they get re-expanded; and -# * put everything else in single quotes, so that it's not re-expanded. + +# Add default JVM options here. You can also use JAVA_OPTS and GRADLE_OPTS to pass JVM options to this script. +DEFAULT_JVM_OPTS='"-Xmx64m" "-Xms64m"' + +# Collect all arguments for the java command: +# * DEFAULT_JVM_OPTS, JAVA_OPTS, JAVA_OPTS, and optsEnvironmentVar are not allowed to contain shell fragments, +# and any embedded shellness will be escaped. +# * For example: A user cannot expect ${Hostname} to be expanded, as it is an environment variable and will be +# treated as '${Hostname}' itself on the command line. set -- \ "-Dorg.gradle.appname=$APP_BASE_NAME" \ @@ -205,6 +217,12 @@ set -- \ org.gradle.wrapper.GradleWrapperMain \ "$@" +# Stop when "xargs" is not available. +if ! command -v xargs >/dev/null 2>&1 +then + die "xargs is not available" +fi + # Use "xargs" to parse quoted args. # # With -n1 it outputs one arg per line, with the quotes and backslashes removed. diff --git a/gradlew.bat b/gradlew.bat index 6a68175..9b42019 100644 --- a/gradlew.bat +++ b/gradlew.bat @@ -1,89 +1,94 @@ -@rem -@rem Copyright 2015 the original author or authors. -@rem -@rem Licensed under the Apache License, Version 2.0 (the "License"); -@rem you may not use this file except in compliance with the License. -@rem You may obtain a copy of the License at -@rem -@rem https://www.apache.org/licenses/LICENSE-2.0 -@rem -@rem Unless required by applicable law or agreed to in writing, software -@rem distributed under the License is distributed on an "AS IS" BASIS, -@rem WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -@rem See the License for the specific language governing permissions and -@rem limitations under the License. -@rem - -@if "%DEBUG%" == "" @echo off -@rem ########################################################################## -@rem -@rem Gradle startup script for Windows -@rem -@rem ########################################################################## - -@rem Set local scope for the variables with windows NT shell -if "%OS%"=="Windows_NT" setlocal - -set DIRNAME=%~dp0 -if "%DIRNAME%" == "" set DIRNAME=. -set APP_BASE_NAME=%~n0 -set APP_HOME=%DIRNAME% - -@rem Resolve any "." and ".." in APP_HOME to make it shorter. -for %%i in ("%APP_HOME%") do set APP_HOME=%%~fi - -@rem Add default JVM options here. You can also use JAVA_OPTS and GRADLE_OPTS to pass JVM options to this script. -set DEFAULT_JVM_OPTS=-Dfile.encoding=UTF-8 "-Xmx64m" "-Xms64m" - -@rem Find java.exe -if defined JAVA_HOME goto findJavaFromJavaHome - -set JAVA_EXE=java.exe -%JAVA_EXE% -version >NUL 2>&1 -if "%ERRORLEVEL%" == "0" goto execute - -echo. -echo ERROR: JAVA_HOME is not set and no 'java' command could be found in your PATH. -echo. -echo Please set the JAVA_HOME variable in your environment to match the -echo location of your Java installation. - -goto fail - -:findJavaFromJavaHome -set JAVA_HOME=%JAVA_HOME:"=% -set JAVA_EXE=%JAVA_HOME%/bin/java.exe - -if exist "%JAVA_EXE%" goto execute - -echo. -echo ERROR: JAVA_HOME is set to an invalid directory: %JAVA_HOME% -echo. -echo Please set the JAVA_HOME variable in your environment to match the -echo location of your Java installation. - -goto fail - -:execute -@rem Setup the command line - -set CLASSPATH=%APP_HOME%\gradle\wrapper\gradle-wrapper.jar - - -@rem Execute Gradle -"%JAVA_EXE%" %DEFAULT_JVM_OPTS% %JAVA_OPTS% %GRADLE_OPTS% "-Dorg.gradle.appname=%APP_BASE_NAME%" -classpath "%CLASSPATH%" org.gradle.wrapper.GradleWrapperMain %* - -:end -@rem End local scope for the variables with windows NT shell -if "%ERRORLEVEL%"=="0" goto mainEnd - -:fail -rem Set variable GRADLE_EXIT_CONSOLE if you need the _script_ return code instead of -rem the _cmd.exe /c_ return code! -if not "" == "%GRADLE_EXIT_CONSOLE%" exit 1 -exit /b 1 - -:mainEnd -if "%OS%"=="Windows_NT" endlocal - -:omega +@rem +@rem Copyright 2015 the original author or authors. +@rem +@rem Licensed under the Apache License, Version 2.0 (the "License"); +@rem you may not use this file except in compliance with the License. +@rem You may obtain a copy of the License at +@rem +@rem https://www.apache.org/licenses/LICENSE-2.0 +@rem +@rem Unless required by applicable law or agreed to in writing, software +@rem distributed under the License is distributed on an "AS IS" BASIS, +@rem WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +@rem See the License for the specific language governing permissions and +@rem limitations under the License. +@rem +@rem SPDX-License-Identifier: Apache-2.0 +@rem + +@if "%DEBUG%"=="" @echo off +@rem ########################################################################## +@rem +@rem Gradle startup script for Windows +@rem +@rem ########################################################################## + +@rem Set local scope for the variables with windows NT shell +if "%OS%"=="Windows_NT" setlocal + +set DIRNAME=%~dp0 +if "%DIRNAME%"=="" set DIRNAME=. +@rem This is normally unused +set APP_BASE_NAME=%~n0 +set APP_HOME=%DIRNAME% + +@rem Resolve any "." and ".." in APP_HOME to make it shorter. +for %%i in ("%APP_HOME%") do set APP_HOME=%%~fi + +@rem Add default JVM options here. You can also use JAVA_OPTS and GRADLE_OPTS to pass JVM options to this script. +set DEFAULT_JVM_OPTS="-Xmx64m" "-Xms64m" + +@rem Find java.exe +if defined JAVA_HOME goto findJavaFromJavaHome + +set JAVA_EXE=java.exe +%JAVA_EXE% -version >NUL 2>&1 +if %ERRORLEVEL% equ 0 goto execute + +echo. 1>&2 +echo ERROR: JAVA_HOME is not set and no 'java' command could be found in your PATH. 1>&2 +echo. 1>&2 +echo Please set the JAVA_HOME variable in your environment to match the 1>&2 +echo location of your Java installation. 1>&2 + +goto fail + +:findJavaFromJavaHome +set JAVA_HOME=%JAVA_HOME:"=% +set JAVA_EXE=%JAVA_HOME%/bin/java.exe + +if exist "%JAVA_EXE%" goto execute + +echo. 1>&2 +echo ERROR: JAVA_HOME is set to an invalid directory: %JAVA_HOME% 1>&2 +echo. 1>&2 +echo Please set the JAVA_HOME variable in your environment to match the 1>&2 +echo location of your Java installation. 1>&2 + +goto fail + +:execute +@rem Setup the command line + +set CLASSPATH=%APP_HOME%\gradle\wrapper\gradle-wrapper.jar + + +@rem Execute Gradle +"%JAVA_EXE%" %DEFAULT_JVM_OPTS% %JAVA_OPTS% %GRADLE_OPTS% "-Dorg.gradle.appname=%APP_BASE_NAME%" -classpath "%CLASSPATH%" org.gradle.wrapper.GradleWrapperMain %* + +:end +@rem End local scope for the variables with windows NT shell +if %ERRORLEVEL% equ 0 goto mainEnd + +:fail +rem Set variable GRADLE_EXIT_CONSOLE if you need the _script_ return code instead of +rem the _cmd.exe /c_ return code! +set EXIT_CODE=%ERRORLEVEL% +if %EXIT_CODE% equ 0 set EXIT_CODE=1 +if not ""=="%GRADLE_EXIT_CONSOLE%" exit %EXIT_CODE% +exit /b %EXIT_CODE% + +:mainEnd +if "%OS%"=="Windows_NT" endlocal + +:omega diff --git a/helpers/personalized_search_ranking_quickstart.sh b/helpers/personalized_search_ranking_quickstart.sh new file mode 100755 index 0000000..c78b3ad --- /dev/null +++ b/helpers/personalized_search_ranking_quickstart.sh @@ -0,0 +1,414 @@ +#!/bin/bash + +set -o errexit +set -o errtrace +set -o pipefail +set -o nounset + +# Some useful constants +readonly DOCKER_IMAGE_TAG="opensearch-with-personalized-search-ranking" +readonly OPENSEARCH_VERSION="2.9.0" + +# +# Set default values for OpenSearch (+Dashboards) image tags and plugin URL. +# We can pass them from outside to override these settings for testing purposes. +# +if [ -z "${OPENSEARCH_IMAGE_TAG:-}" ]; then + OPENSEARCH_IMAGE_TAG="opensearchproject/opensearch:${OPENSEARCH_VERSION}" +fi +if [ -z "${OPENSEARCH_DASHBOARDS_IMAGE_TAG:-}" ]; then + OPENSEARCH_DASHBOARDS_IMAGE_TAG="opensearchproject/opensearch-dashboards:${OPENSEARCH_VERSION}" +fi +if [ -z "${SEARCH_PROCESSOR_PLUGIN_URL:-}" ]; then + SEARCH_PROCESSOR_PLUGIN_URL="https://github.com/opensearch-project/search-processor/releases/download/${OPENSEARCH_VERSION}/opensearch-search-processor-${OPENSEARCH_VERSION}.0.zip" +fi + +function print_help() { + cat << EOF +Usage: $0 [-r ] [--profile ] + [--volume-name ] [--admin-password ] + -r | --region The AWS region for the Personalize Intelligent Ranking + service endpoint. If not specified, will read from the + AWS CLI for the default profile. + --profile The AWS profile to use for credentials. If not set, then + the script will try first to use credentials from the + environment, then from the default AWS profile. + --volume-name Without this option, the OpenSearch container will write + the index to ephemeral container storage, which is lost when + the container is removed. Using this option will map the + named Docker volume to \$OPENSEARCH_ROOT/data, so index data + will persist across executions. If the named volume does not + exist, it will be created. + --admin-password For OpenSearch 2.12 and higher, we no longer use a default + password of "admin" for the admin user. Instead, the value + passed to this parameter will be used as the admin password. + For OpenSearch versions prior to 2.12, this argument will be + ignored with a warning. + + NOTE: If the --profile option is not specified, the script will attempt to read AWS + credentials (access/secret key, optional session token) from environment variables, + and then from the default AWS profile, in order to pass them to the OpenSearch keystore + to be used to connect to the Personalize Intelligent ranking service. + If no credentials are found, the script WILL NOT pass credentials to the OpenSearch + keystore. When running a reranking request, the ranking plugin may rely on + instance profile credentials delivered through the EC2 metadata service, or credentials + from the ECS metadata service. +EOF +} + +# +# Parse and validate arguments +# + +while [ "$#" -gt 0 ]; do + case $1 in + -r | --region ) + shift + AWS_REGION=$1 + shift + ;; + -h | --help ) + print_help + exit 0 + ;; + --profile ) + shift + AWS_PROFILE=$1 + shift + ;; + --volume-name ) + shift + VOLUME_NAME=$1 + shift + ;; + --admin-password ) + shift + OPENSEARCH_INITIAL_ADMIN_PASSWORD="$1" + shift + ;; + esac +done + +# Starting in 2.12.0, security demo configuration script requires an initial admin password +OPENSEARCH_REQUIRED_VERSION="2.12.0" +COMPARE_VERSION=`echo $OPENSEARCH_REQUIRED_VERSION $OPENSEARCH_VERSION | tr ' ' '\n' | sort -V | uniq | head -n 1` +if [ "$COMPARE_VERSION" != "$OPENSEARCH_REQUIRED_VERSION" ]; then + if [ -n "${OPENSEARCH_INITIAL_ADMIN_PASSWORD:-}" ]; then + echo "WARNING: The --admin-password setting has no effect on OpenSearch ${OPENSEARCH_VERSION}. The admin password will be 'admin'." + fi + OPENSEARCH_INITIAL_ADMIN_PASSWORD="admin" +elif [ -z "${OPENSEARCH_INITIAL_ADMIN_PASSWORD:-}" ]; then + echo "Starting with OpenSearch 2.12, you must specify the admin password with the --admin-password parameter." + exit 1 +fi + +# +# Determine which credentials and region to use. By the end of this block, all specified +# credentials will be loaded into environment variables (or we fail with an explanatory +# error message). +# +if [ -n "${AWS_PROFILE:-}" ]; then + # Load everything from the specified profile + AWS_ACCESS_KEY_ID=$(aws --profile ${AWS_PROFILE} configure get aws_access_key_id || echo) + AWS_SECRET_ACCESS_KEY=$(aws --profile ${AWS_PROFILE} configure get aws_secret_access_key || echo) + AWS_SESSION_TOKEN=$(aws --profile ${AWS_PROFILE} configure get aws_session_token || echo) + if [ -z "${AWS_ACCESS_KEY_ID:-}" ] || [ -z "${AWS_SECRET_ACCESS_KEY:-}" ]; then + >&2 echo "Unable to load credentials from profile ${AWS_PROFILE}" + exit 1 + elif [ -z "${AWS_SESSION_TOKEN}" ]; then + echo "Using AWS credentials (aws_access_key_id and aws_secret_access_key) from profile ${AWS_PROFILE}" + else + echo "Using AWS credentials (aws_access_key_id, aws_secret_access_key, and aws_session_token) from profile ${AWS_PROFILE}" + fi + if [ -z "${AWS_REGION:-}" ]; then + AWS_REGION=$(aws --profile ${AWS_PROFILE} configure get region || echo) + if [ -n "${AWS_REGION:-}" ]; then + echo "Using AWS region ${AWS_REGION} from profile ${AWS_PROFILE}" + else + >&2 echo "Argument [-r | --region] not specified and unable to infer region from profile ${AWS_PROFILE}" + exit 1 + fi + fi +else + # No profile set + if [ -z "${AWS_ACCESS_KEY_ID:-}" ]; then + if [ -z "${AWS_SECRET_ACCESS_KEY:-}" ]; then + echo "No profile set and no credentials in environment. Trying to load default profile credentials." + AWS_ACCESS_KEY_ID=$(aws configure get aws_access_key_id || echo) + AWS_SECRET_ACCESS_KEY=$(aws configure get aws_secret_access_key || echo) + AWS_SESSION_TOKEN=$(aws configure get aws_session_token || echo) + if [ -z "${AWS_ACCESS_KEY_ID:-}" ] || [ -z "${AWS_SECRET_ACCESS_KEY}" ]; then + echo "Unable to load credentials from default profile. No credentials will be passed to the OpenSearch keystore." + echo "OpenSearch will use the default credential provider chain to access Personalize, which may rely on EC2 instance" + echo "profile credentials or credentials from ECS metadata service." + elif [ -z "${AWS_SESSION_TOKEN}" ]; then + echo "Using AWS credentials (aws_access_key_id and aws_secret_access_key) from default profile." + else + echo "Using AWS credentials (aws_access_key_id, aws_secret_access_key, and aws_session_token) from default profile." + fi + else + >&2 echo "Environment variable AWS_SECRET_ACCESS_KEY is specified, but AWS_ACCESS_KEY_ID is not." + >&2 echo "Unable to determine which credentials to use." + exit 1 + fi + else + # AWS_ACCCESS_KEY_ID is set + if [ -z "${AWS_SECRET_ACCESS_KEY:-}" ]; then + >&2 echo "Environment variable AWS_ACCESS_KEY_ID is specified, but AWS_SECRET_ACCESS_KEY is not." + >&2 echo "Unable to determine which credentials to use." + exit 1 + else + if [ -n "${AWS_SESSION_TOKEN:-}" ]; then + echo "Using credentials from environment (AWS_ACCESS_KEY_ID, AWS_SECRET_ACCESS_KEY, and AWS_SESSION_TOKEN)." + else + echo "Using credentials from environment (AWS_ACCESS_KEY_ID and AWS_SECRET_ACCESS_KEY)." + fi + fi + fi + if [ -z "${AWS_REGION:-}" ]; then + AWS_REGION=$(aws configure get region || echo) + if [ -n "${AWS_REGION:-}" ]; then + echo "Using AWS region ${AWS_REGION} from default profile" + else + >&2 echo "Argument [-r | --region] not specified and unable to infer region from default profile" + exit 1 + fi + fi +fi + +echo "Established AWS key id, secret key, session token and region." + +# +# Create a unique directory to hold the Dockerfile and docker-compose.yml files. +# +PLATFORM=$(uname) +if [ "${PLATFORM}" == "Darwin" ]; then + DOCKER_BUILD_DIR=$(mktemp -d opensearch-personalize-intelligent-ranking-docker.XXXX) +else + # Assume GNU mktemp + DOCKER_BUILD_DIR=$(mktemp -d -p . opensearch-personalize-intelligent-ranking-docker.XXXX) +fi + +cd ${DOCKER_BUILD_DIR} + +SUFFIX=$(echo ${DOCKER_BUILD_DIR} | sed s/.*\.//) +echo $SUFFIX + +echo "Running in $(pwd)" + +# +# Construct a Dockerfile that installs the search-processor plugin in the target image +# +cat >Dockerfile <>personalized_search_ranking.credentials <>personalized_search_ranking.credentials <install_credentials.sh <<"EOF" +#!/bin/bash + +for l in $(cat $1); do + KEY=$(echo $l | cut -f1 -d:) + VALUE=$(echo $l | cut -f2 -d:) + echo $VALUE | /usr/share/opensearch/bin/opensearch-keystore add $KEY --stdin +done +EOF + chmod 755 install_credentials.sh + cat >>Dockerfile << EOF + +# Push credentials to keystore +COPY --chown=opensearch:opensearch install_credentials.sh /tmp +COPY --chown=opensearch:opensearch personalized_search_ranking.credentials /tmp +RUN /usr/share/opensearch/bin/opensearch-keystore create +RUN --mount=type=secret,id=credentials,target=/tmp/personalized_search_ranking.credentials,required=true,mode=0444 \ + /tmp/install_credentials.sh /tmp/personalized_search_ranking.credentials +EOF +fi +echo "Opensearch credentials saved" +# +# Build and tag the Docker image with the plugin (and maybe credentials in the keystore) +# +if [ -f personalized_search_ranking.credentials ]; then + DOCKER_BUILDKIT=1 docker build --tag ${DOCKER_IMAGE_TAG} --secret id=credentials,src=personalized_search_ranking.credentials . + rm personalized_search_ranking.credentials +else + docker build --tag ${DOCKER_IMAGE_TAG} . +fi +echo "Docker image built and tagged with credentials" +# +# Make sure we have opensearch-dashboards: +# +docker pull ${OPENSEARCH_DASHBOARDS_IMAGE_TAG} +echo "Docker image pulled" + +if [ -n "${VOLUME_NAME:-}" ]; then + if ! docker volume inspect ${VOLUME_NAME}> /dev/null; then + echo "Creating volume ${VOLUME_NAME}"; + docker volume create ${VOLUME_NAME} + fi + DATA_DIR_BLOCK=" volumes: + - ${VOLUME_NAME}:/usr/share/opensearch/data" + VOLUME_BLOCK="volumes: + ${VOLUME_NAME}: + external: true" +fi +echo "Volume created" + + + +# +# Create a docker-compose.yml file that will launch an OpenSearch node with the image we +# just built and an OpenSearch Dashboards node that points to the OpenSearch node. +# +cat >docker-compose.yml <cleanup_resources.sh <README <" https://localhost:9200/ + +Index some data on OpenSearch by following instructions at +https://opensearch.org/docs/latest/opensearch/index-data/ + + +Connect to OpenSearch Dashboards with a web browser at http://localhost:5601/, +using username admin and password admin. Select "Search Relevance" from the +top-left menu. In the resulting UI, you can submit a query without Personalized +search ranking and one with Personalized search Ranking. + +To configure and setup Personalize search ranking, run a curl command as follows: + +curl -X PUT "https://localhost:9200/_search/pipeline/intelligent_ranking" -u 'admin:' --insecure -H 'Content-Type: application/json' -d' +{ + "description": "A pipeline to apply custom reranking", + "response_processors" : [ + { + "personalized_search_ranking" : { + "campaign_arn" : "", + "item_id_field" : "", + "recipe" : "", + "weight" : "", + "iam_role_arn": "", + "aws_region": "" + } + } + ] +}' + +Interact with the Docker containers using docker-compose from directory + $(pwd) + +Some helpful docker-compose commands: + + docker-compose logs opensearch-node + Outputs latest logs from the OpenSearch server + + docker-compose logs opensearch-dashboard + Outputs latest logs from the OpenSearch Dashboard server + + docker-compose down + Shut down and clean up both containers. + + docker-compose up -d + Bring both containers back up. + +You can clean up all Docker containers and any execution plans created (if +applicable) by running + $(pwd)/cleanup_resources.sh + +The full text of this message is also available at + $(pwd)/README +EOF +cat README diff --git a/helpers/search_processing_kendra_quickstart.sh b/helpers/search_processing_kendra_quickstart.sh index 09a90c8..3c0bba4 100755 --- a/helpers/search_processing_kendra_quickstart.sh +++ b/helpers/search_processing_kendra_quickstart.sh @@ -7,7 +7,7 @@ set -o nounset # Some useful constants readonly DOCKER_IMAGE_TAG="opensearch-with-ranking-plugin" -readonly OPENSEARCH_VERSION="2.6.0" +readonly OPENSEARCH_VERSION="2.7.0" # # Set default values for OpenSearch (+Dashboards) image tags and plugin URL. @@ -27,7 +27,7 @@ function print_help() { cat << EOF Usage: $0 [-p ] [-r ] [-e ] [--profile ] [--create-execution-plan] - [--volume-name ] + [--volume-name ] [--admin-password ] -p | --execution-plan-id The ID returned from Kendra Intelligent Ranking service from the call to CreateRescoreExecutionPlan. Required if --create-execution-plan is not set. @@ -50,6 +50,11 @@ Usage: $0 [-p ] [-r ] [-e ] named Docker volume to \$OPENSEARCH_ROOT/data, so index data will persist across executions. If the named volume does not exist, it will be created. + --admin-password For OpenSearch 2.12 and higher, we no longer use a default + password of "admin" for the admin user. Instead, the value + passed to this parameter will be used as the admin password. + For OpenSearch versions prior to 2.12, this argument will be + ignored with a warning. NOTE: If the --profile option is not specified, the script will attempt to read AWS credentials (access/secret key, optional session token) from environment variables, @@ -101,6 +106,11 @@ while [ "$#" -gt 0 ]; do VOLUME_NAME=$1 shift ;; + --admin-password ) + shift + OPENSEARCH_INITIAL_ADMIN_PASSWORD="$1" + shift + ;; esac done @@ -121,6 +131,19 @@ if [ "${FAILED_VALIDATION}" == "1" ]; then exit 1 fi +# Starting in 2.12.0, security demo configuration script requires an initial admin password +OPENSEARCH_REQUIRED_VERSION="2.12.0" +COMPARE_VERSION=`echo $OPENSEARCH_REQUIRED_VERSION $OPENSEARCH_VERSION | tr ' ' '\n' | sort -V | uniq | head -n 1` +if [ "$COMPARE_VERSION" != "$OPENSEARCH_REQUIRED_VERSION" ]; then + if [ -n "${OPENSEARCH_INITIAL_ADMIN_PASSWORD:-}" ]; then + echo "WARNING: The --admin-password setting has no effect on OpenSearch ${OPENSEARCH_VERSION}. The admin password will be 'admin'." + fi + OPENSEARCH_INITIAL_ADMIN_PASSWORD="admin" +elif [ -z "${OPENSEARCH_INITIAL_ADMIN_PASSWORD:-}" ]; then + echo "Starting with OpenSearch 2.12, you must specify the admin password with the --admin-password parameter." + exit 1 +fi + # # Determine which credentials and region to use. By the end of this block, all specified # credentials will be loaded into environment variables (or we fail with an explanatory @@ -379,6 +402,7 @@ services: - kendra_intelligent_ranking.service.endpoint=${KENDRA_RANKING_ENDPOINT} - kendra_intelligent_ranking.service.region=${AWS_REGION} - kendra_intelligent_ranking.service.execution_plan_id=${EXECUTION_PLAN_ID} + - OPENSEARCH_INITIAL_ADMIN_PASSWORD=${OPENSEARCH_INITIAL_ADMIN_PASSWORD} ulimits: memlock: soft: -1 @@ -446,8 +470,8 @@ cat >README <" https://localhost:9200/ Index some data on OpenSearch by following instructions at https://opensearch.org/docs/latest/opensearch/index-data/ diff --git a/release-notes/search-processor.release-notes-2.7.0.md b/release-notes/search-processor.release-notes-2.7.0.md new file mode 100644 index 0000000..8ba8b47 --- /dev/null +++ b/release-notes/search-processor.release-notes-2.7.0.md @@ -0,0 +1,11 @@ +## Version 2.7.0 Release Notes + +### Enhancements +* Improve test coverage [#99](https://github.com/opensearch-project/search-processor/pull/99) + +### Infrastructure +* Updating build.gradle to use snapshot by default [#124](https://github.com/opensearch-project/search-processor/pull/124) +* Accommodate changes to XContent classes [#125](https://github.com/opensearch-project/search-processor/pull/125) + +### Documentation +* Prepping for 2.7.0 release [#130](https://github.com/opensearch-project/search-processor/pull/130) diff --git a/release-notes/search-processor.release-notes-2.8.0.md b/release-notes/search-processor.release-notes-2.8.0.md new file mode 100644 index 0000000..287dbce --- /dev/null +++ b/release-notes/search-processor.release-notes-2.8.0.md @@ -0,0 +1,9 @@ + +## Version 2.8.0.0 Release Notes + +Compatible with OpenSearch 2.8.0 + + +### Features +* Add KendraRankingResponseProcessor ([#137](https://github.com/opensearch-project/search-processor/pull/137)) +* Personalized intelligent ranking for open search requests ([#138](https://github.com/opensearch-project/search-processor/pull/138)) \ No newline at end of file diff --git a/release-notes/search-processor.release-notes-2.9.0.md b/release-notes/search-processor.release-notes-2.9.0.md new file mode 100644 index 0000000..8dd944d --- /dev/null +++ b/release-notes/search-processor.release-notes-2.9.0.md @@ -0,0 +1,9 @@ +## Version 2.9.0.0 Release Notes + +Compatible with OpenSearch 2.9.0 + + +### Enhancements +* Support contextual metadata to use when getting personalized reranking ([#144](https://github.com/opensearch-project/search-processor/pull/144)) +* Scoring logic for reranking open search hits based on personalize campaign response ([#147](https://github.com/opensearch-project/search-processor/pull/147)) +* Bring Search Pipeline Update and Add IgnoreFailure and PipelineContext To Each Processors ([#152](https://github.com/opensearch-project/search-processor/pull/152)) diff --git a/settings.gradle b/settings.gradle index cb5e308..96ba866 100644 --- a/settings.gradle +++ b/settings.gradle @@ -8,3 +8,5 @@ */ rootProject.name = 'search-processor' +include 'amazon-kendra-intelligent-ranking' +include 'amazon-personalize-ranking' diff --git a/src/main/java/org/opensearch/search/relevance/SearchRelevancePlugin.java b/src/main/java/org/opensearch/search/relevance/SearchRelevancePlugin.java deleted file mode 100644 index b2143f9..0000000 --- a/src/main/java/org/opensearch/search/relevance/SearchRelevancePlugin.java +++ /dev/null @@ -1,110 +0,0 @@ -/* - * SPDX-License-Identifier: Apache-2.0 - * - * The OpenSearch Contributors require contributions made to - * this file be licensed under the Apache-2.0 license or a - * compatible open source license. - */ -package org.opensearch.search.relevance; - -import java.util.ArrayList; -import java.util.Arrays; -import java.util.Collection; -import java.util.Collections; -import java.util.List; -import java.util.Map; -import java.util.function.Supplier; -import java.util.stream.Collectors; - -import org.opensearch.action.support.ActionFilter; -import org.opensearch.client.Client; -import org.opensearch.cluster.metadata.IndexNameExpressionResolver; -import org.opensearch.cluster.service.ClusterService; -import org.opensearch.common.io.stream.NamedWriteableRegistry; -import org.opensearch.common.settings.Setting; -import org.opensearch.core.xcontent.NamedXContentRegistry; -import org.opensearch.env.Environment; -import org.opensearch.env.NodeEnvironment; -import org.opensearch.plugins.ActionPlugin; -import org.opensearch.plugins.Plugin; -import org.opensearch.plugins.SearchPlugin; -import org.opensearch.repositories.RepositoriesService; -import org.opensearch.script.ScriptService; -import org.opensearch.search.relevance.actionfilter.SearchActionFilter; -import org.opensearch.search.relevance.client.OpenSearchClient; -import org.opensearch.search.relevance.configuration.ResultTransformerConfigurationFactory; -import org.opensearch.search.relevance.transformer.kendraintelligentranking.client.KendraClientSettings; -import org.opensearch.search.relevance.transformer.kendraintelligentranking.client.KendraHttpClient; -import org.opensearch.search.relevance.configuration.SearchConfigurationExtBuilder; -import org.opensearch.search.relevance.transformer.kendraintelligentranking.KendraIntelligentRanker; -import org.opensearch.search.relevance.transformer.ResultTransformer; -import org.opensearch.search.relevance.transformer.kendraintelligentranking.configuration.KendraIntelligentRankerSettings; -import org.opensearch.search.relevance.transformer.kendraintelligentranking.configuration.KendraIntelligentRankingConfigurationFactory; -import org.opensearch.threadpool.ThreadPool; -import org.opensearch.watcher.ResourceWatcherService; - -public class SearchRelevancePlugin extends Plugin implements ActionPlugin, SearchPlugin { - - private OpenSearchClient openSearchClient; - private KendraHttpClient kendraClient; - private KendraIntelligentRanker kendraIntelligentRanker; - - private Collection getAllResultTransformers() { - // Initialize and add other transformers here - return List.of(this.kendraIntelligentRanker); - } - - private Collection getResultTransformerConfigurationFactories() { - return List.of(KendraIntelligentRankingConfigurationFactory.INSTANCE); - } - - @Override - public List getActionFilters() { - return Arrays.asList(new SearchActionFilter(getAllResultTransformers(), openSearchClient)); - } - - @Override - public List> getSettings() { - // NOTE: cannot use kendraIntelligentRanker.getTransformerSettings because the object is not yet created - List> allTransformerSettings = new ArrayList<>(); - allTransformerSettings.addAll(KendraIntelligentRankerSettings.getAllSettings()); - // Add settings for other transformers here - return allTransformerSettings; - } - - @Override - public Collection createComponents( - Client client, - ClusterService clusterService, - ThreadPool threadPool, - ResourceWatcherService resourceWatcherService, - ScriptService scriptService, - NamedXContentRegistry xContentRegistry, - Environment environment, - NodeEnvironment nodeEnvironment, - NamedWriteableRegistry namedWriteableRegistry, - IndexNameExpressionResolver indexNameExpressionResolver, - Supplier repositoriesServiceSupplier - ) { - this.openSearchClient = new OpenSearchClient(client); - this.kendraClient = new KendraHttpClient(KendraClientSettings.getClientSettings(environment.settings())); - this.kendraIntelligentRanker = new KendraIntelligentRanker(this.kendraClient); - - return Arrays.asList( - this.openSearchClient, - this.kendraClient, - this.kendraIntelligentRanker - ); - } - - @Override - public List> getSearchExts() { - Map resultTransformerMap = getResultTransformerConfigurationFactories().stream() - .collect(Collectors.toMap(ResultTransformerConfigurationFactory::getName, i -> i)); - return Collections.singletonList( - new SearchExtSpec<>(SearchConfigurationExtBuilder.NAME, - input -> new SearchConfigurationExtBuilder(input, resultTransformerMap), - parser -> SearchConfigurationExtBuilder.parse(parser, resultTransformerMap))); - } - -} diff --git a/src/main/java/org/opensearch/search/relevance/configuration/ConfigurationUtils.java b/src/main/java/org/opensearch/search/relevance/configuration/ConfigurationUtils.java deleted file mode 100644 index ed7b60d..0000000 --- a/src/main/java/org/opensearch/search/relevance/configuration/ConfigurationUtils.java +++ /dev/null @@ -1,95 +0,0 @@ -/* - * SPDX-License-Identifier: Apache-2.0 - * - * The OpenSearch Contributors require contributions made to - * this file be licensed under the Apache-2.0 license or a - * compatible open source license. - */ -package org.opensearch.search.relevance.configuration; - -import org.opensearch.action.search.SearchRequest; -import org.opensearch.common.settings.Settings; -import org.opensearch.search.SearchExtBuilder; -import org.opensearch.search.relevance.transformer.ResultTransformer; - -import java.util.ArrayList; -import java.util.Comparator; -import java.util.List; -import java.util.Map; -import java.util.stream.Collectors; - -import static org.opensearch.search.relevance.configuration.Constants.RESULT_TRANSFORMER_SETTING_PREFIX; - -public class ConfigurationUtils { - - /** - * Get result transformer configurations from Search Request - * @param settings all index settings configured for this plugin - * @return ordered and validated list of result transformers, empty list if not specified - */ - public static List getResultTransformersFromIndexConfiguration(Settings settings, - Map resultTransformerMap) { - List indexLevelConfigs = new ArrayList<>(); - - if (settings != null) { - if (settings.getGroups(RESULT_TRANSFORMER_SETTING_PREFIX) != null) { - for (Map.Entry transformerSettings : settings.getGroups(RESULT_TRANSFORMER_SETTING_PREFIX).entrySet()) { - if (resultTransformerMap.containsKey(transformerSettings.getKey())) { - ResultTransformer transformer = resultTransformerMap.get(transformerSettings.getKey()); - indexLevelConfigs.add(transformer.getConfigurationFactory().configure(transformerSettings.getValue())); - } - } - } - } - - return reorderAndValidateConfigs(indexLevelConfigs); - } - - /** - * Get result transformer configurations from Search Request - * @param searchRequest input request - * @return ordered and validated list of result transformers, empty list if not specified - */ - public static List getResultTransformersFromRequestConfiguration( - final SearchRequest searchRequest) { - - // Fetch result transformers specified in request - SearchConfigurationExtBuilder requestLevelSearchConfiguration = null; - if (searchRequest.source() != null && searchRequest.source().ext() != null && !searchRequest.source().ext().isEmpty()) { - // Filter ext builders by name - List extBuilders = searchRequest.source().ext().stream() - .filter(searchExtBuilder -> SearchConfigurationExtBuilder.NAME.equals(searchExtBuilder.getWriteableName())) - .collect(Collectors.toList()); - if (!extBuilders.isEmpty()) { - requestLevelSearchConfiguration = (SearchConfigurationExtBuilder) extBuilders.get(0); - } - } - - List requestLevelConfigs = new ArrayList<>(); - if (requestLevelSearchConfiguration != null) { - requestLevelConfigs = reorderAndValidateConfigs(requestLevelSearchConfiguration.getResultTransformers()); - } - return requestLevelConfigs; - } - - /** - * Sort configurations in ascending order of invocation, and validate - * @param configs list of result transformer configurations - * @return ordered and validated list of result transformers - */ - public static List reorderAndValidateConfigs( - final List configs) throws IllegalArgumentException { - - // Sort - configs.sort(Comparator.comparingInt(ResultTransformerConfiguration::getOrder)); - - for (int i = 0; i < configs.size(); ++i) { - if (configs.get(i).getOrder() != (i + 1)) { - throw new IllegalArgumentException("Expected order [" + (i + 1) + "] for transformer [" + - configs.get(i).getTransformerName() + "], but found [" + configs.get(i).getOrder() + "]"); - } - } - - return configs; - } -} diff --git a/src/yamlRestTest/resources/rest-api-spec/test/10_basic.yml b/src/yamlRestTest/resources/rest-api-spec/test/10_basic.yml deleted file mode 100644 index 9898f9f..0000000 --- a/src/yamlRestTest/resources/rest-api-spec/test/10_basic.yml +++ /dev/null @@ -1,8 +0,0 @@ -"Test that the plugin is loaded in OpenSearch": - - do: - cat.plugins: - local: true - h: component - - - match: - $body: /^opensearch-search-processor-\d+.\d+.\d+.\d+\n$/