From 3c8f47fc9dcb0f02cd50a96c8b79a5bba1e44cc1 Mon Sep 17 00:00:00 2001 From: wooh Date: Thu, 30 Jul 2026 03:32:50 +0900 Subject: [PATCH 1/2] =?UTF-8?q?[Fix]=20Corpus=20=EC=9E=84=EB=B2=A0?= =?UTF-8?q?=EB=94=A9=20Cohere=20=ED=81=B4=EB=9D=BC=EC=9D=B4=EC=96=B8?= =?UTF-8?q?=ED=8A=B8=20=EC=A4=91=EB=B3=B5=20=EC=A0=9C=EA=B1=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - CohereCorpusEmbeddingClient를 global CohereEmbeddingClient 위임 구조로 변경 - cohere.api-key/base-url/embedding 설정으로 Cohere 설정 통합 - corpus embedding 모델 참조를 CohereProperties 기준으로 정리 - 기존 pgvector 저장 및 similarity search 구조는 유지 - Cohere corpus 위임 테스트 보강 --- .../AnalysisInputFingerprintProvider.java | 5 +- .../service/CohereCorpusEmbeddingClient.java | 103 +----------------- .../service/CorpusEmbeddingSyncService.java | 9 +- .../MockQuestionCacheVersionProvider.java | 5 +- .../resources/application-analysis-eval.yaml | 2 - src/main/resources/application-dev.yaml | 5 - src/main/resources/application-prod.yaml | 5 - .../CohereCorpusEmbeddingClientTest.java | 45 ++++++++ ...ockQuestionCachePropertiesTestSupport.java | 8 +- .../MockQuestionCacheVersionProviderTest.java | 25 +++-- 10 files changed, 84 insertions(+), 128 deletions(-) create mode 100644 src/test/java/com/jobdri/jobdri_api/domain/corpus/service/CohereCorpusEmbeddingClientTest.java diff --git a/src/main/java/com/jobdri/jobdri_api/domain/analysis/service/core/AnalysisInputFingerprintProvider.java b/src/main/java/com/jobdri/jobdri_api/domain/analysis/service/core/AnalysisInputFingerprintProvider.java index ef483ca3..c38917c1 100644 --- a/src/main/java/com/jobdri/jobdri_api/domain/analysis/service/core/AnalysisInputFingerprintProvider.java +++ b/src/main/java/com/jobdri/jobdri_api/domain/analysis/service/core/AnalysisInputFingerprintProvider.java @@ -9,6 +9,7 @@ import com.jobdri.jobdri_api.domain.analysis.service.ai.FewShotPromptProvider; import com.jobdri.jobdri_api.domain.corpus.service.CorpusRetrievalService; import com.jobdri.jobdri_api.domain.jobposting.entity.JobPosting; +import com.jobdri.jobdri_api.global.cohere.CohereProperties; import org.springframework.beans.factory.annotation.Value; import org.springframework.stereotype.Component; import org.springframework.util.StringUtils; @@ -40,10 +41,10 @@ public class AnalysisInputFingerprintProvider { public AnalysisInputFingerprintProvider( ObjectMapper objectMapper, FewShotPromptProvider fewShotPromptProvider, + CohereProperties cohereProperties, @Value("${openai.model.cover-letter-analysis:gpt-4o-mini}") String analysisModel, @Value("${analysis.two-pass.enabled:false}") boolean twoPassEnabled, @Value("${analysis.mode:}") String analysisMode, - @Value("${app.corpus.embedding.model:embed-v4.0}") String embeddingModel, @Value("${app.analysis.retrieval.jd-limit:3}") int jdLimit, @Value("${app.analysis.retrieval.question-limit:5}") int questionLimit ) { @@ -52,7 +53,7 @@ public AnalysisInputFingerprintProvider( this.analysisModel = analysisModel; this.twoPassEnabled = twoPassEnabled; this.analysisMode = analysisMode; - this.embeddingModel = embeddingModel; + this.embeddingModel = cohereProperties.embedding().model(); this.jdLimit = jdLimit; this.questionLimit = questionLimit; } diff --git a/src/main/java/com/jobdri/jobdri_api/domain/corpus/service/CohereCorpusEmbeddingClient.java b/src/main/java/com/jobdri/jobdri_api/domain/corpus/service/CohereCorpusEmbeddingClient.java index 85d151d7..86dbe902 100644 --- a/src/main/java/com/jobdri/jobdri_api/domain/corpus/service/CohereCorpusEmbeddingClient.java +++ b/src/main/java/com/jobdri/jobdri_api/domain/corpus/service/CohereCorpusEmbeddingClient.java @@ -1,115 +1,22 @@ package com.jobdri.jobdri_api.domain.corpus.service; -import com.fasterxml.jackson.databind.JsonNode; -import com.fasterxml.jackson.databind.ObjectMapper; +import com.jobdri.jobdri_api.global.cohere.CohereEmbeddingClient; import lombok.RequiredArgsConstructor; -import org.springframework.beans.factory.annotation.Value; -import org.springframework.http.HttpHeaders; -import org.springframework.http.MediaType; -import org.springframework.http.client.SimpleClientHttpRequestFactory; import org.springframework.stereotype.Component; -import org.springframework.util.StringUtils; -import org.springframework.web.client.RestClient; -import java.time.Duration; import java.util.List; @Component @RequiredArgsConstructor public class CohereCorpusEmbeddingClient implements CorpusEmbeddingClient { - private final RestClient.Builder restClientBuilder; - private final ObjectMapper objectMapper; - - @Value("${cohere.api.key:}") - private String cohereApiKey; - - @Value("${app.corpus.embedding.model:embed-v4.0}") - private String embeddingModel; - - @Value("${app.corpus.embedding.output-dimension:1024}") - private int outputDimension; + private final CohereEmbeddingClient cohereEmbeddingClient; @Override public List embed(List texts, InputType inputType) { - if (!StringUtils.hasText(cohereApiKey)) { - throw new IllegalStateException("Cohere API 키가 설정되지 않았습니다."); - } - if (texts == null || texts.isEmpty()) { - return List.of(); - } - - SimpleClientHttpRequestFactory requestFactory = new SimpleClientHttpRequestFactory(); - requestFactory.setConnectTimeout(Duration.ofSeconds(5)); - requestFactory.setReadTimeout(Duration.ofSeconds(10)); - - RestClient client = restClientBuilder - .baseUrl("https://api.cohere.com") - .requestFactory(requestFactory) - .defaultHeader(HttpHeaders.AUTHORIZATION, "Bearer " + cohereApiKey) - .defaultHeader(HttpHeaders.CONTENT_TYPE, MediaType.APPLICATION_JSON_VALUE) - .build(); - - String responseBody = client.post() - .uri("/v2/embed") - .body(new EmbedRequest( - texts, - embeddingModel, - inputType.value(), - outputDimension, - List.of("float") - )) - .retrieve() - .body(String.class); - - return parseEmbeddings(responseBody); - } - - private float[] toFloatArray(List values) { - float[] array = new float[values.size()]; - for (int i = 0; i < values.size(); i++) { - array[i] = values.get(i).floatValue(); + if (inputType == InputType.SEARCH_QUERY) { + return List.of(cohereEmbeddingClient.embedQuery(texts == null || texts.isEmpty() ? null : texts.getFirst())); } - return array; + return cohereEmbeddingClient.embedDocuments(texts); } - - private List parseEmbeddings(String responseBody) { - if (!StringUtils.hasText(responseBody)) { - throw new IllegalStateException("Cohere 임베딩 응답이 비어 있습니다."); - } - - try { - JsonNode root = objectMapper.readTree(responseBody); - JsonNode floatEmbeddings = root.path("embeddings").path("float"); - if (!floatEmbeddings.isArray()) { - throw new IllegalStateException("Cohere 임베딩 응답 형식이 예상과 다릅니다."); - } - - List result = new java.util.ArrayList<>(); - for (JsonNode embeddingNode : floatEmbeddings) { - if (!embeddingNode.isArray()) { - throw new IllegalStateException("Cohere 임베딩 벡터 형식이 예상과 다릅니다."); - } - - float[] vector = new float[embeddingNode.size()]; - for (int i = 0; i < embeddingNode.size(); i++) { - vector[i] = embeddingNode.get(i).floatValue(); - } - result.add(vector); - } - return result; - } catch (Exception e) { - throw new IllegalStateException("Cohere 임베딩 응답 파싱에 실패했습니다.", e); - } - } - - private record EmbedRequest( - List texts, - String model, - String input_type, - Integer output_dimension, - List embedding_types - ) { - } - } diff --git a/src/main/java/com/jobdri/jobdri_api/domain/corpus/service/CorpusEmbeddingSyncService.java b/src/main/java/com/jobdri/jobdri_api/domain/corpus/service/CorpusEmbeddingSyncService.java index 082bf9fb..dd1c7dc8 100644 --- a/src/main/java/com/jobdri/jobdri_api/domain/corpus/service/CorpusEmbeddingSyncService.java +++ b/src/main/java/com/jobdri/jobdri_api/domain/corpus/service/CorpusEmbeddingSyncService.java @@ -5,6 +5,7 @@ import com.jobdri.jobdri_api.domain.corpus.entity.MockQuestionCorpus; import com.jobdri.jobdri_api.domain.corpus.repository.MockJobPostingCorpusRepository; import com.jobdri.jobdri_api.domain.corpus.repository.MockQuestionCorpusRepository; +import com.jobdri.jobdri_api.global.cohere.CohereProperties; import com.pgvector.PGvector; import lombok.RequiredArgsConstructor; import org.springframework.beans.factory.annotation.Value; @@ -46,12 +47,10 @@ ON CONFLICT (corpus_id) updated_at = EXCLUDED.updated_at """; - @Value("${app.corpus.embedding.model:embed-v4.0}") - private String embeddingModel; - @Value("${app.corpus.embedding.batch-size:32}") private int batchSize; + private final CohereProperties cohereProperties; private final MockJobPostingCorpusRepository mockJobPostingCorpusRepository; private final MockQuestionCorpusRepository mockQuestionCorpusRepository; private final CorpusEmbeddingClient corpusEmbeddingClient; @@ -61,7 +60,7 @@ ON CONFLICT (corpus_id) public CorpusEmbeddingSyncResponse syncAll(Integer limit) { int jobPostingCount = syncJobPostingEmbeddings(limit); int questionCount = syncQuestionEmbeddings(limit); - return new CorpusEmbeddingSyncResponse(jobPostingCount, questionCount, embeddingModel); + return new CorpusEmbeddingSyncResponse(jobPostingCount, questionCount, cohereProperties.embedding().model()); } @Transactional @@ -116,7 +115,7 @@ private void upsertVectors(String sql, List ids, List embeddings) Timestamp now = Timestamp.valueOf(LocalDateTime.now()); for (int i = 0; i < ids.size(); i++) { statement.setLong(1, ids.get(i)); - statement.setString(2, embeddingModel); + statement.setString(2, cohereProperties.embedding().model()); statement.setObject(3, new PGvector(embeddings.get(i))); statement.setTimestamp(4, now); statement.setTimestamp(5, now); diff --git a/src/main/java/com/jobdri/jobdri_api/domain/jobposting/service/MockQuestionCacheVersionProvider.java b/src/main/java/com/jobdri/jobdri_api/domain/jobposting/service/MockQuestionCacheVersionProvider.java index f8629f1e..0c8b7ee6 100644 --- a/src/main/java/com/jobdri/jobdri_api/domain/jobposting/service/MockQuestionCacheVersionProvider.java +++ b/src/main/java/com/jobdri/jobdri_api/domain/jobposting/service/MockQuestionCacheVersionProvider.java @@ -1,6 +1,7 @@ package com.jobdri.jobdri_api.domain.jobposting.service; import com.jobdri.jobdri_api.domain.corpus.service.CorpusRetrievalService; +import com.jobdri.jobdri_api.global.cohere.CohereProperties; import org.springframework.beans.factory.annotation.Value; import org.springframework.stereotype.Component; import org.springframework.util.StringUtils; @@ -24,14 +25,14 @@ public class MockQuestionCacheVersionProvider { public MockQuestionCacheVersionProvider( MockQuestionCacheProperties mockQuestionCacheProperties, + CohereProperties cohereProperties, @Value("${openai.model.job-posting-extractor:gpt-4o-mini}") String extractionModel, - @Value("${app.corpus.embedding.model:embed-v4.0}") String embeddingModel, @Value("${app.analysis.retrieval.jd-limit:3}") int jdLimit, @Value("${app.analysis.retrieval.question-limit:5}") int questionLimit ) { this.mockQuestionCacheProperties = mockQuestionCacheProperties; this.extractionModel = extractionModel; - this.embeddingModel = embeddingModel; + this.embeddingModel = cohereProperties.embedding().model(); this.jdLimit = jdLimit; this.questionLimit = questionLimit; } diff --git a/src/main/resources/application-analysis-eval.yaml b/src/main/resources/application-analysis-eval.yaml index a23869e0..9101fc91 100644 --- a/src/main/resources/application-analysis-eval.yaml +++ b/src/main/resources/application-analysis-eval.yaml @@ -93,8 +93,6 @@ cohere: dimension: ${COHERE_EMBEDDING_DIMENSION:1024} connect-timeout: ${COHERE_EMBEDDING_CONNECT_TIMEOUT:3s} read-timeout: ${COHERE_EMBEDDING_READ_TIMEOUT:15s} - api: - key: ${COHERE_API_KEY:} jwt: secret: diff --git a/src/main/resources/application-dev.yaml b/src/main/resources/application-dev.yaml index f025533a..0dc9f988 100644 --- a/src/main/resources/application-dev.yaml +++ b/src/main/resources/application-dev.yaml @@ -116,9 +116,6 @@ app: allowed-root: ${APP_CORPUS_IMPORT_ALLOWED_ROOT:} embedding: sync-on-startup: ${APP_CORPUS_EMBEDDING_SYNC_ON_STARTUP:false} - model: ${APP_CORPUS_EMBEDDING_MODEL:embed-v4.0} - output-dimension: ${APP_CORPUS_EMBEDDING_OUTPUT_DIMENSION:1024} - document-input-type: ${APP_CORPUS_EMBEDDING_DOCUMENT_INPUT_TYPE:search_document} batch-size: ${APP_CORPUS_EMBEDDING_BATCH_SIZE:32} analysis: retrieval: @@ -170,8 +167,6 @@ cohere: dimension: ${COHERE_EMBEDDING_DIMENSION:1024} connect-timeout: ${COHERE_EMBEDDING_CONNECT_TIMEOUT:3s} read-timeout: ${COHERE_EMBEDDING_READ_TIMEOUT:15s} - api: - key: ${COHERE_API_KEY:} payment: coupon: diff --git a/src/main/resources/application-prod.yaml b/src/main/resources/application-prod.yaml index 7bda61fa..5ad569e0 100644 --- a/src/main/resources/application-prod.yaml +++ b/src/main/resources/application-prod.yaml @@ -116,9 +116,6 @@ app: allowed-root: ${APP_CORPUS_IMPORT_ALLOWED_ROOT:} embedding: sync-on-startup: ${APP_CORPUS_EMBEDDING_SYNC_ON_STARTUP:false} - model: ${APP_CORPUS_EMBEDDING_MODEL:embed-v4.0} - output-dimension: ${APP_CORPUS_EMBEDDING_OUTPUT_DIMENSION:1024} - document-input-type: ${APP_CORPUS_EMBEDDING_DOCUMENT_INPUT_TYPE:search_document} batch-size: ${APP_CORPUS_EMBEDDING_BATCH_SIZE:32} analysis: retrieval: @@ -171,8 +168,6 @@ cohere: dimension: ${COHERE_EMBEDDING_DIMENSION:1024} connect-timeout: ${COHERE_EMBEDDING_CONNECT_TIMEOUT:3s} read-timeout: ${COHERE_EMBEDDING_READ_TIMEOUT:15s} - api: - key: ${COHERE_API_KEY:} payment: coupon: diff --git a/src/test/java/com/jobdri/jobdri_api/domain/corpus/service/CohereCorpusEmbeddingClientTest.java b/src/test/java/com/jobdri/jobdri_api/domain/corpus/service/CohereCorpusEmbeddingClientTest.java new file mode 100644 index 00000000..a2497632 --- /dev/null +++ b/src/test/java/com/jobdri/jobdri_api/domain/corpus/service/CohereCorpusEmbeddingClientTest.java @@ -0,0 +1,45 @@ +package com.jobdri.jobdri_api.domain.corpus.service; + +import com.jobdri.jobdri_api.global.cohere.CohereEmbeddingClient; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +import java.util.List; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +class CohereCorpusEmbeddingClientTest { + + @Test + @DisplayName("문서 임베딩은 전역 CohereEmbeddingClient의 embedDocuments에 위임한다") + void delegatesDocumentEmbeddingToGlobalClient() { + CohereEmbeddingClient globalClient = mock(CohereEmbeddingClient.class); + List texts = List.of("문서 1", "문서 2"); + List embeddings = List.of(new float[]{1.0f}, new float[]{2.0f}); + when(globalClient.embedDocuments(texts)).thenReturn(embeddings); + + CohereCorpusEmbeddingClient client = new CohereCorpusEmbeddingClient(globalClient); + + assertThat(client.embed(texts, CorpusEmbeddingClient.InputType.SEARCH_DOCUMENT)) + .isSameAs(embeddings); + verify(globalClient).embedDocuments(texts); + } + + @Test + @DisplayName("검색 쿼리 임베딩은 전역 CohereEmbeddingClient의 embedQuery에 위임한다") + void delegatesQueryEmbeddingToGlobalClient() { + CohereEmbeddingClient globalClient = mock(CohereEmbeddingClient.class); + when(globalClient.embedQuery("검색 질의")).thenReturn(new float[]{1.0f, 2.0f}); + + CohereCorpusEmbeddingClient client = new CohereCorpusEmbeddingClient(globalClient); + + List result = client.embed(List.of("검색 질의"), CorpusEmbeddingClient.InputType.SEARCH_QUERY); + + assertThat(result).hasSize(1); + assertThat(result.getFirst()).containsExactly(1.0f, 2.0f); + verify(globalClient).embedQuery("검색 질의"); + } +} diff --git a/src/test/java/com/jobdri/jobdri_api/domain/jobposting/service/MockQuestionCachePropertiesTestSupport.java b/src/test/java/com/jobdri/jobdri_api/domain/jobposting/service/MockQuestionCachePropertiesTestSupport.java index f03409d2..8129d0a5 100644 --- a/src/test/java/com/jobdri/jobdri_api/domain/jobposting/service/MockQuestionCachePropertiesTestSupport.java +++ b/src/test/java/com/jobdri/jobdri_api/domain/jobposting/service/MockQuestionCachePropertiesTestSupport.java @@ -1,5 +1,7 @@ package com.jobdri.jobdri_api.domain.jobposting.service; +import com.jobdri.jobdri_api.global.cohere.CohereProperties; + final class MockQuestionCachePropertiesTestSupport { static final String VERSION_PREFIX = "v1"; @@ -21,8 +23,12 @@ static MockQuestionCacheProperties createProperties() { static MockQuestionCacheVersionProvider createVersionProvider() { return new MockQuestionCacheVersionProvider( createProperties(), + new CohereProperties( + "test-api-key", + "https://api.cohere.com", + new CohereProperties.Embedding("embed-v4.0", 1024, null, null) + ), "gpt-4o-mini", - "embed-v4.0", 3, 5 ); diff --git a/src/test/java/com/jobdri/jobdri_api/domain/jobposting/service/MockQuestionCacheVersionProviderTest.java b/src/test/java/com/jobdri/jobdri_api/domain/jobposting/service/MockQuestionCacheVersionProviderTest.java index a0cd8cc0..6f4110bd 100644 --- a/src/test/java/com/jobdri/jobdri_api/domain/jobposting/service/MockQuestionCacheVersionProviderTest.java +++ b/src/test/java/com/jobdri/jobdri_api/domain/jobposting/service/MockQuestionCacheVersionProviderTest.java @@ -1,5 +1,6 @@ package com.jobdri.jobdri_api.domain.jobposting.service; +import com.jobdri.jobdri_api.global.cohere.CohereProperties; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; @@ -24,15 +25,15 @@ void currentVersionChangesWhenModelChanges() { MockQuestionCacheProperties properties = MockQuestionCachePropertiesTestSupport.createProperties(); MockQuestionCacheVersionProvider baseline = new MockQuestionCacheVersionProvider( properties, + cohereProperties("embed-v4.0"), "gpt-4o-mini", - "embed-v4.0", 3, 5 ); MockQuestionCacheVersionProvider changedModel = new MockQuestionCacheVersionProvider( properties, + cohereProperties("embed-v4.0"), "gpt-5-mini", - "embed-v4.0", 3, 5 ); @@ -46,15 +47,15 @@ void currentVersionChangesWhenEmbeddingModelChanges() { MockQuestionCacheProperties properties = MockQuestionCachePropertiesTestSupport.createProperties(); MockQuestionCacheVersionProvider baseline = new MockQuestionCacheVersionProvider( properties, + cohereProperties("embed-v4.0"), "gpt-4o-mini", - "embed-v4.0", 3, 5 ); MockQuestionCacheVersionProvider changedEmbeddingModel = new MockQuestionCacheVersionProvider( properties, + cohereProperties("embed-v5.0"), "gpt-4o-mini", - "embed-v5.0", 3, 5 ); @@ -68,15 +69,15 @@ void currentVersionChangesWhenJdLimitChanges() { MockQuestionCacheProperties properties = MockQuestionCachePropertiesTestSupport.createProperties(); MockQuestionCacheVersionProvider baseline = new MockQuestionCacheVersionProvider( properties, + cohereProperties("embed-v4.0"), "gpt-4o-mini", - "embed-v4.0", 3, 5 ); MockQuestionCacheVersionProvider changedJdLimit = new MockQuestionCacheVersionProvider( properties, + cohereProperties("embed-v4.0"), "gpt-4o-mini", - "embed-v4.0", 4, 5 ); @@ -90,19 +91,27 @@ void currentVersionChangesWhenQuestionLimitChanges() { MockQuestionCacheProperties properties = MockQuestionCachePropertiesTestSupport.createProperties(); MockQuestionCacheVersionProvider baseline = new MockQuestionCacheVersionProvider( properties, + cohereProperties("embed-v4.0"), "gpt-4o-mini", - "embed-v4.0", 3, 5 ); MockQuestionCacheVersionProvider changedQuestionLimit = new MockQuestionCacheVersionProvider( properties, + cohereProperties("embed-v4.0"), "gpt-4o-mini", - "embed-v4.0", 3, 6 ); assertThat(changedQuestionLimit.currentVersion()).isNotEqualTo(baseline.currentVersion()); } + + private CohereProperties cohereProperties(String model) { + return new CohereProperties( + "test-api-key", + "https://api.cohere.com", + new CohereProperties.Embedding(model, 1024, null, null) + ); + } } From 7067a884bcf652f0e3f87ddda6621fc9de4f4cf6 Mon Sep 17 00:00:00 2001 From: wooh Date: Thu, 30 Jul 2026 03:39:31 +0900 Subject: [PATCH 2/2] =?UTF-8?q?[Fix]=20Cohere=20=EC=9E=84=EB=B2=A0?= =?UTF-8?q?=EB=94=A9=20=ED=81=B4=EB=9D=BC=EC=9D=B4=EC=96=B8=ED=8A=B8=20?= =?UTF-8?q?=ED=92=80=EB=A7=81=20=EB=B0=8F=20=EC=9E=AC=EC=8B=9C=EB=8F=84=20?= =?UTF-8?q?=EC=A0=81=EC=9A=A9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - CohereEmbeddingClient에 HttpClient5 connection pool 적용 - 429/5xx 응답에 bounded exponential backoff 재시도 추가 - Retry-After 헤더 기반 대기 처리 - readTimeout 예외 매핑 회귀 테스트 추가 --- build.gradle | 1 + .../global/cohere/CohereEmbeddingClient.java | 130 +++++++++++++++++- .../cohere/CohereEmbeddingClientTest.java | 119 +++++++++++++++- 3 files changed, 241 insertions(+), 9 deletions(-) diff --git a/build.gradle b/build.gradle index 040c4c2b..fb46b242 100644 --- a/build.gradle +++ b/build.gradle @@ -41,6 +41,7 @@ dependencies { implementation 'org.springframework.boot:spring-boot-starter-jdbc' implementation 'org.springframework.boot:spring-boot-starter-validation' implementation 'org.springframework.boot:spring-boot-starter-web' + implementation 'org.apache.httpcomponents.client5:httpclient5' implementation 'org.apache.poi:poi-ooxml:5.4.1' implementation 'com.pgvector:pgvector:0.1.6' diff --git a/src/main/java/com/jobdri/jobdri_api/global/cohere/CohereEmbeddingClient.java b/src/main/java/com/jobdri/jobdri_api/global/cohere/CohereEmbeddingClient.java index 6ea425d6..42d4b6de 100644 --- a/src/main/java/com/jobdri/jobdri_api/global/cohere/CohereEmbeddingClient.java +++ b/src/main/java/com/jobdri/jobdri_api/global/cohere/CohereEmbeddingClient.java @@ -5,15 +5,28 @@ import com.jobdri.jobdri_api.global.cohere.dto.CohereEmbeddingRequest; import com.jobdri.jobdri_api.global.cohere.dto.CohereEmbeddingResponse; import lombok.extern.slf4j.Slf4j; +import org.apache.hc.client5.http.config.RequestConfig; +import org.apache.hc.client5.http.impl.classic.CloseableHttpClient; +import org.apache.hc.client5.http.impl.classic.HttpClients; +import org.apache.hc.client5.http.impl.io.PoolingHttpClientConnectionManager; +import org.apache.hc.client5.http.impl.io.PoolingHttpClientConnectionManagerBuilder; +import org.apache.hc.core5.util.Timeout; import org.springframework.http.HttpHeaders; import org.springframework.http.MediaType; -import org.springframework.http.client.SimpleClientHttpRequestFactory; +import org.springframework.http.client.ClientHttpResponse; +import org.springframework.http.client.HttpComponentsClientHttpRequestFactory; import org.springframework.stereotype.Component; import org.springframework.util.StringUtils; import org.springframework.web.client.ResourceAccessException; import org.springframework.web.client.RestClient; import org.springframework.web.client.RestClientException; +import java.io.IOException; +import java.time.Duration; +import java.time.OffsetDateTime; +import java.time.ZonedDateTime; +import java.time.format.DateTimeFormatter; +import java.time.format.DateTimeParseException; import java.util.ArrayList; import java.util.List; @@ -21,6 +34,11 @@ @Slf4j public class CohereEmbeddingClient { private static final int MAX_TEXTS_PER_REQUEST = 96; + private static final int MAX_TOTAL_CONNECTIONS = 100; + private static final int MAX_CONNECTIONS_PER_ROUTE = 20; + private static final int MAX_TRANSIENT_ATTEMPTS = 3; + private static final Duration INITIAL_RETRY_BACKOFF = Duration.ofMillis(200); + private static final Duration MAX_RETRY_BACKOFF = Duration.ofSeconds(2); private static final String INPUT_TYPE_SEARCH_DOCUMENT = "search_document"; private static final String INPUT_TYPE_SEARCH_QUERY = "search_query"; private static final List FLOAT_EMBEDDING_TYPE = List.of("float"); @@ -68,6 +86,30 @@ private List embed(List texts, String inputType) { } private CohereEmbeddingResponse callCohere(CohereEmbeddingRequest request) { + Duration backoff = INITIAL_RETRY_BACKOFF; + for (int attempt = 1; attempt <= MAX_TRANSIENT_ATTEMPTS; attempt++) { + try { + return callCohereOnce(request); + } catch (TransientCohereException e) { + if (attempt == MAX_TRANSIENT_ATTEMPTS) { + throw unavailable("Cohere Embed API가 일시적으로 응답할 수 없습니다.", e); + } + Duration delay = e.retryAfter() != null ? e.retryAfter() : backoff; + log.warn( + "Cohere Embed API transient failure. attempt={}, maxAttempts={}, retryAfterMs={}, message={}", + attempt, + MAX_TRANSIENT_ATTEMPTS, + delay.toMillis(), + e.getMessage() + ); + sleepBeforeRetry(delay); + backoff = nextBackoff(backoff); + } + } + throw unavailable("Cohere Embed API가 일시적으로 응답할 수 없습니다."); + } + + private CohereEmbeddingResponse callCohereOnce(CohereEmbeddingRequest request) { try { return restClient.post() .uri("/v2/embed") @@ -76,8 +118,11 @@ private CohereEmbeddingResponse callCohere(CohereEmbeddingRequest request) { .retrieve() .onStatus( status -> status.value() == 429 || status.is5xxServerError(), - (ignoredRequest, ignoredResponse) -> { - throw unavailable("Cohere Embed API가 일시적으로 응답할 수 없습니다."); + (ignoredRequest, response) -> { + throw new TransientCohereException( + "Cohere Embed API transient status=" + response.getStatusCode().value(), + retryAfter(response) + ); } ) .onStatus( @@ -161,13 +206,73 @@ private List validateTexts(List texts) { return List.copyOf(normalizedTexts); } - private static SimpleClientHttpRequestFactory requestFactory(CohereProperties properties) { - SimpleClientHttpRequestFactory requestFactory = new SimpleClientHttpRequestFactory(); - requestFactory.setConnectTimeout(properties.embedding().connectTimeout()); + private static HttpComponentsClientHttpRequestFactory requestFactory(CohereProperties properties) { + RequestConfig requestConfig = RequestConfig.custom() + .setConnectTimeout(timeout(properties.embedding().connectTimeout())) + .setResponseTimeout(timeout(properties.embedding().readTimeout())) + .build(); + PoolingHttpClientConnectionManager connectionManager = PoolingHttpClientConnectionManagerBuilder.create() + .setMaxConnTotal(MAX_TOTAL_CONNECTIONS) + .setMaxConnPerRoute(MAX_CONNECTIONS_PER_ROUTE) + .build(); + CloseableHttpClient httpClient = HttpClients.custom() + .setConnectionManager(connectionManager) + .setDefaultRequestConfig(requestConfig) + .build(); + HttpComponentsClientHttpRequestFactory requestFactory = new HttpComponentsClientHttpRequestFactory(httpClient); + requestFactory.setConnectionRequestTimeout(properties.embedding().connectTimeout()); requestFactory.setReadTimeout(properties.embedding().readTimeout()); return requestFactory; } + private static Timeout timeout(Duration duration) { + return Timeout.ofMilliseconds(duration.toMillis()); + } + + private static Duration retryAfter(ClientHttpResponse response) throws IOException { + String value = response.getHeaders().getFirst(HttpHeaders.RETRY_AFTER); + if (!StringUtils.hasText(value)) { + return null; + } + try { + long seconds = Long.parseLong(value.trim()); + return seconds <= 0 ? Duration.ZERO : Duration.ofSeconds(seconds); + } catch (NumberFormatException ignored) { + try { + Duration duration = Duration.between(OffsetDateTime.now(), OffsetDateTime.parse(value.trim())); + return duration.isNegative() ? Duration.ZERO : duration; + } catch (DateTimeParseException ignoredDate) { + try { + Duration duration = Duration.between( + ZonedDateTime.now(), + ZonedDateTime.parse(value.trim(), DateTimeFormatter.RFC_1123_DATE_TIME) + ); + return duration.isNegative() ? Duration.ZERO : duration; + } catch (DateTimeParseException ignoredHttpDate) { + return null; + } + } + } + } + + private static Duration nextBackoff(Duration current) { + Duration next = current.multipliedBy(2); + return next.compareTo(MAX_RETRY_BACKOFF) > 0 ? MAX_RETRY_BACKOFF : next; + } + + private static void sleepBeforeRetry(Duration delay) { + try { + Thread.sleep(delay.toMillis()); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + throw new GeneralException( + GeneralErrorCode.SERVICE_UNAVAILABLE, + "Cohere Embed API 재시도 대기 중 인터럽트되었습니다.", + e + ); + } + } + private GeneralException invalidParameter(String message) { return new GeneralException(GeneralErrorCode.INVALID_PARAMETER, message); } @@ -179,4 +284,17 @@ private GeneralException unavailable(String message) { private GeneralException unavailable(String message, Throwable cause) { return new GeneralException(GeneralErrorCode.SERVICE_UNAVAILABLE, message, cause); } + + private static final class TransientCohereException extends RuntimeException { + private final Duration retryAfter; + + private TransientCohereException(String message, Duration retryAfter) { + super(message); + this.retryAfter = retryAfter; + } + + private Duration retryAfter() { + return retryAfter; + } + } } diff --git a/src/test/java/com/jobdri/jobdri_api/global/cohere/CohereEmbeddingClientTest.java b/src/test/java/com/jobdri/jobdri_api/global/cohere/CohereEmbeddingClientTest.java index 6cc106c5..51970f26 100644 --- a/src/test/java/com/jobdri/jobdri_api/global/cohere/CohereEmbeddingClientTest.java +++ b/src/test/java/com/jobdri/jobdri_api/global/cohere/CohereEmbeddingClientTest.java @@ -15,6 +15,7 @@ import java.time.Duration; import java.util.ArrayList; import java.util.List; +import java.util.concurrent.atomic.AtomicInteger; import java.util.concurrent.atomic.AtomicReference; import static org.assertj.core.api.Assertions.assertThat; @@ -118,6 +119,25 @@ void transientCohereErrors() throws Exception { } } + @Test + @DisplayName("Cohere 429와 5xx는 bounded retry 후 성공 응답을 반환한다") + void retryTransientCohereError() throws Exception { + AtomicReference requestJson = new AtomicReference<>(); + try (TestCohereServer server = startTransientThenSuccessServer( + 429, + responseJson(3, 1), + requestJson + )) { + CohereEmbeddingClient client = client(server.baseUrl(), "test-api-key", 3); + + float[] embedding = client.embedQuery("Spring Boot 기반 REST API 개발"); + + assertThat(embedding).hasSize(3); + assertThat(server.requestCount()).isEqualTo(2); + assertThat(requestJson.get().get("input_type").asText()).isEqualTo("search_query"); + } + } + @Test @DisplayName("Cohere 400, 401, 403은 요청 또는 설정 오류로 변환한다") void requestOrConfigurationErrors() throws Exception { @@ -132,6 +152,28 @@ void requestOrConfigurationErrors() throws Exception { } } + @Test + @DisplayName("Cohere 응답이 readTimeout보다 지연되면 timeout 예외로 변환한다") + void readTimeout() throws Exception { + try (TestCohereServer server = startServer( + 200, + responseJson(3, 1), + new AtomicReference<>(), + Duration.ofMillis(500) + )) { + CohereEmbeddingClient client = client( + server.baseUrl(), + "test-api-key", + 3, + Duration.ofMillis(100) + ); + + assertThatThrownBy(() -> client.embedQuery("query")) + .isInstanceOfSatisfying(GeneralException.class, exception -> + assertThat(exception.getCode()).isEqualTo(GeneralErrorCode.EXTERNAL_SERVICE_TIMEOUT)); + } + } + @Test @DisplayName("응답 embedding 개수가 요청 texts 개수와 다르면 예외 처리한다") void mismatchedEmbeddingCount() throws Exception { @@ -178,6 +220,10 @@ void emptyResponseBodyOrEmbeddings() throws Exception { } private CohereEmbeddingClient client(String baseUrl, String apiKey, int dimension) { + return client(baseUrl, apiKey, dimension, Duration.ofSeconds(2)); + } + + private CohereEmbeddingClient client(String baseUrl, String apiKey, int dimension, Duration readTimeout) { return new CohereEmbeddingClient( new CohereProperties( apiKey, @@ -186,7 +232,7 @@ private CohereEmbeddingClient client(String baseUrl, String apiKey, int dimensio "embed-v4.0", dimension, Duration.ofSeconds(1), - Duration.ofSeconds(2) + readTimeout ) ), RestClient.builder() @@ -212,15 +258,29 @@ private TestCohereServer startServer( int status, String responseBody, AtomicReference requestJson + ) throws IOException { + return startServer(status, responseBody, requestJson, Duration.ZERO); + } + + private TestCohereServer startServer( + int status, + String responseBody, + AtomicReference requestJson, + Duration responseDelay ) throws IOException { HttpServer server = HttpServer.create(new InetSocketAddress(0), 0); AtomicReference authorizationHeader = new AtomicReference<>(); + AtomicInteger requestCount = new AtomicInteger(); server.createContext("/v2/embed", exchange -> { + requestCount.incrementAndGet(); authorizationHeader.set(exchange.getRequestHeaders().getFirst("Authorization")); String requestBody = new String(exchange.getRequestBody().readAllBytes(), StandardCharsets.UTF_8); if (!requestBody.isBlank()) { requestJson.set(objectMapper.readTree(requestBody)); } + if (!responseDelay.isZero()) { + sleep(responseDelay); + } byte[] body = responseBody.getBytes(StandardCharsets.UTF_8); exchange.getResponseHeaders().set("Content-Type", "application/json"); exchange.sendResponseHeaders(status, body.length); @@ -228,16 +288,65 @@ private TestCohereServer startServer( exchange.close(); }); server.start(); - return new TestCohereServer(server, authorizationHeader); + return new TestCohereServer(server, authorizationHeader, requestCount); + } + + private TestCohereServer startTransientThenSuccessServer( + int transientStatus, + String successResponseBody, + AtomicReference requestJson + ) throws IOException { + HttpServer server = HttpServer.create(new InetSocketAddress(0), 0); + AtomicReference authorizationHeader = new AtomicReference<>(); + AtomicInteger requestCount = new AtomicInteger(); + server.createContext("/v2/embed", exchange -> { + int requestNumber = requestCount.incrementAndGet(); + authorizationHeader.set(exchange.getRequestHeaders().getFirst("Authorization")); + String requestBody = new String(exchange.getRequestBody().readAllBytes(), StandardCharsets.UTF_8); + if (!requestBody.isBlank()) { + requestJson.set(objectMapper.readTree(requestBody)); + } + if (requestNumber == 1) { + byte[] body = "{\"message\":\"temporary\"}".getBytes(StandardCharsets.UTF_8); + exchange.getResponseHeaders().set("Retry-After", "0"); + exchange.getResponseHeaders().set("Content-Type", "application/json"); + exchange.sendResponseHeaders(transientStatus, body.length); + exchange.getResponseBody().write(body); + exchange.close(); + return; + } + byte[] body = successResponseBody.getBytes(StandardCharsets.UTF_8); + exchange.getResponseHeaders().set("Content-Type", "application/json"); + exchange.sendResponseHeaders(200, body.length); + exchange.getResponseBody().write(body); + exchange.close(); + }); + server.start(); + return new TestCohereServer(server, authorizationHeader, requestCount); + } + + private static void sleep(Duration duration) { + try { + Thread.sleep(duration.toMillis()); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + throw new RuntimeException(e); + } } private static final class TestCohereServer implements AutoCloseable { private final HttpServer server; private final AtomicReference authorizationHeader; + private final AtomicInteger requestCount; - private TestCohereServer(HttpServer server, AtomicReference authorizationHeader) { + private TestCohereServer( + HttpServer server, + AtomicReference authorizationHeader, + AtomicInteger requestCount + ) { this.server = server; this.authorizationHeader = authorizationHeader; + this.requestCount = requestCount; } String baseUrl() { @@ -248,6 +357,10 @@ String authorizationHeader() { return authorizationHeader.get(); } + int requestCount() { + return requestCount.get(); + } + @Override public void close() { server.stop(0);