feat(personal): add vision OCR gateway
This commit is contained in:
+203
@@ -0,0 +1,203 @@
|
||||
package org.dromara.aihr.personal.service;
|
||||
|
||||
import com.fasterxml.jackson.databind.JsonNode;
|
||||
import com.fasterxml.jackson.databind.ObjectMapper;
|
||||
import com.fasterxml.jackson.databind.node.ArrayNode;
|
||||
import com.fasterxml.jackson.databind.node.ObjectNode;
|
||||
import org.springframework.beans.factory.annotation.Value;
|
||||
import org.springframework.dao.DataAccessException;
|
||||
import org.springframework.jdbc.core.JdbcTemplate;
|
||||
import org.springframework.stereotype.Service;
|
||||
|
||||
import java.io.InputStream;
|
||||
import java.net.URI;
|
||||
import java.net.http.HttpClient;
|
||||
import java.net.http.HttpRequest;
|
||||
import java.net.http.HttpResponse;
|
||||
import java.nio.charset.StandardCharsets;
|
||||
import java.time.Duration;
|
||||
import java.util.Base64;
|
||||
import java.util.List;
|
||||
import java.util.Optional;
|
||||
import java.util.stream.Collectors;
|
||||
|
||||
@Service
|
||||
public class PersonalVisionOcrService {
|
||||
|
||||
private static final String TENANT_ID = "000000";
|
||||
private static final int MAX_RESPONSE_BYTES = 1024 * 1024;
|
||||
private static final String PROMPT =
|
||||
"忠实提取第%d页全部可见文字,保留标题、段落和表格行顺序;不要总结、解释或补写。无可识别文字时返回空字符串。";
|
||||
|
||||
private final ObjectMapper objectMapper;
|
||||
private final boolean aiEnabled;
|
||||
private final boolean chatEnabled;
|
||||
private final RuntimeProvider runtimeProvider;
|
||||
private final VisionCaller caller;
|
||||
|
||||
public PersonalVisionOcrService(ObjectMapper objectMapper, JdbcTemplate jdbcTemplate,
|
||||
@Value("${aihr.ai-runtime.enabled:${AIHR_AI_RUNTIME_ENABLED:true}}")
|
||||
boolean aiEnabled,
|
||||
@Value("${aihr.ai-runtime.chat-enabled:${AIHR_AI_CHAT_ENABLED:true}}")
|
||||
boolean chatEnabled) {
|
||||
this(objectMapper, aiEnabled, chatEnabled, () -> resolveRuntime(jdbcTemplate),
|
||||
PersonalVisionOcrService::callProvider);
|
||||
}
|
||||
|
||||
private PersonalVisionOcrService(ObjectMapper objectMapper, boolean aiEnabled, boolean chatEnabled,
|
||||
RuntimeProvider runtimeProvider, VisionCaller caller) {
|
||||
this.objectMapper = objectMapper;
|
||||
this.aiEnabled = aiEnabled;
|
||||
this.chatEnabled = chatEnabled;
|
||||
this.runtimeProvider = runtimeProvider;
|
||||
this.caller = caller;
|
||||
}
|
||||
|
||||
public static PersonalVisionOcrService forTest(ObjectMapper objectMapper, boolean aiEnabled,
|
||||
boolean chatEnabled, RuntimeProvider runtimeProvider,
|
||||
VisionCaller caller) {
|
||||
return new PersonalVisionOcrService(objectMapper, aiEnabled, chatEnabled, runtimeProvider, caller);
|
||||
}
|
||||
|
||||
public String recognize(byte[] imageBytes, String mimeType, int pageNumber) {
|
||||
if (!aiEnabled || !chatEnabled) {
|
||||
throw new OcrUnavailableException("PERSONAL_OCR_RUNTIME_DISABLED");
|
||||
}
|
||||
if (imageBytes == null || imageBytes.length == 0 || mimeType == null || !mimeType.startsWith("image/")) {
|
||||
throw new IllegalArgumentException("invalid OCR page image");
|
||||
}
|
||||
VisionRuntime runtime = runtimeProvider.resolve()
|
||||
.orElseThrow(() -> new OcrUnavailableException("PERSONAL_OCR_MODEL_UNAVAILABLE"));
|
||||
try {
|
||||
VisionResponse response = caller.send(runtime, requestBody(runtime, imageBytes, mimeType, pageNumber));
|
||||
if (response.statusCode() < 200 || response.statusCode() >= 300) {
|
||||
throw new OcrUnavailableException("PERSONAL_OCR_PROVIDER_FAILED");
|
||||
}
|
||||
JsonNode content = objectMapper.readTree(response.body()).path("choices").path(0)
|
||||
.path("message").path("content");
|
||||
if (!content.isTextual()) {
|
||||
throw new OcrUnavailableException("PERSONAL_OCR_RESPONSE_INVALID");
|
||||
}
|
||||
return normalize(content.asText());
|
||||
} catch (OcrUnavailableException exception) {
|
||||
throw exception;
|
||||
} catch (Exception exception) {
|
||||
throw new OcrUnavailableException("PERSONAL_OCR_PROVIDER_FAILED", exception);
|
||||
}
|
||||
}
|
||||
|
||||
private String requestBody(VisionRuntime runtime, byte[] imageBytes, String mimeType, int pageNumber)
|
||||
throws Exception {
|
||||
ObjectNode body = objectMapper.createObjectNode();
|
||||
body.put("model", runtime.modelName());
|
||||
body.put("temperature", 0);
|
||||
body.put("max_tokens", 4096);
|
||||
ArrayNode messages = body.putArray("messages");
|
||||
ObjectNode user = messages.addObject();
|
||||
user.put("role", "user");
|
||||
ArrayNode content = user.putArray("content");
|
||||
content.addObject().put("type", "text").put("text", PROMPT.formatted(pageNumber));
|
||||
String dataUrl = "data:" + mimeType + ";base64," + Base64.getEncoder().encodeToString(imageBytes);
|
||||
content.addObject().put("type", "image_url").putObject("image_url")
|
||||
.put("url", dataUrl).put("detail", "high");
|
||||
return objectMapper.writeValueAsString(body);
|
||||
}
|
||||
|
||||
private static Optional<VisionRuntime> resolveRuntime(JdbcTemplate jdbcTemplate) {
|
||||
try {
|
||||
List<VisionRuntime> rows = jdbcTemplate.query("""
|
||||
select c.model_name,
|
||||
coalesce(nullif(c.api_host, ''), nullif(p.api_host, '')) resolved_api_host,
|
||||
coalesce(nullif(c.api_key, ''), nullif(p.api_key, '')) resolved_api_key
|
||||
from aihr_model_config c
|
||||
left join aihr_model_provider p
|
||||
on p.tenant_id = c.tenant_id and p.provider_code = c.provider_code
|
||||
where c.tenant_id = ? and c.category in ('vision', 'chat') and c.enabled = 1
|
||||
and (p.status is null or p.status = '0')
|
||||
order by case c.category when 'vision' then 0 else 1 end,
|
||||
case when c.model_show = 'Y' then 0 else 1 end, c.id
|
||||
limit 1
|
||||
""", (rs, rowNum) -> new VisionRuntime(rs.getString("model_name"),
|
||||
rs.getString("resolved_api_host"), rs.getString("resolved_api_key")), TENANT_ID);
|
||||
return rows.stream().filter(runtime -> notBlank(runtime.modelName()) && notBlank(runtime.baseUrl()))
|
||||
.findFirst();
|
||||
} catch (DataAccessException exception) {
|
||||
return Optional.empty();
|
||||
}
|
||||
}
|
||||
|
||||
private static VisionResponse callProvider(VisionRuntime runtime, String body) throws Exception {
|
||||
HttpRequest.Builder request = HttpRequest.newBuilder()
|
||||
.uri(URI.create(normalizeBaseUrl(runtime.baseUrl()) + "/chat/completions"))
|
||||
.timeout(Duration.ofSeconds(120))
|
||||
.header("Content-Type", "application/json")
|
||||
.POST(HttpRequest.BodyPublishers.ofString(body));
|
||||
if (notBlank(runtime.apiKey())) {
|
||||
request.header("Authorization", "Bearer " + runtime.apiKey());
|
||||
}
|
||||
HttpResponse<InputStream> response = HttpClient.newBuilder()
|
||||
.connectTimeout(Duration.ofSeconds(15))
|
||||
.build()
|
||||
.send(request.build(), HttpResponse.BodyHandlers.ofInputStream());
|
||||
try (InputStream input = response.body()) {
|
||||
byte[] bytes = input.readNBytes(MAX_RESPONSE_BYTES + 1);
|
||||
if (bytes.length > MAX_RESPONSE_BYTES) {
|
||||
throw new OcrUnavailableException("PERSONAL_OCR_RESPONSE_TOO_LARGE");
|
||||
}
|
||||
return new VisionResponse(response.statusCode(), new String(bytes, StandardCharsets.UTF_8));
|
||||
}
|
||||
}
|
||||
|
||||
private static String normalize(String value) {
|
||||
if (value == null || value.isBlank()) {
|
||||
return "";
|
||||
}
|
||||
return value.lines().map(String::trim).filter(line -> !line.isBlank()).collect(Collectors.joining("\n"));
|
||||
}
|
||||
|
||||
private static String normalizeBaseUrl(String value) {
|
||||
String normalized = value == null ? "" : value.trim();
|
||||
while (normalized.endsWith("/")) {
|
||||
normalized = normalized.substring(0, normalized.length() - 1);
|
||||
}
|
||||
return normalized.endsWith("/v1") ? normalized : normalized + "/v1";
|
||||
}
|
||||
|
||||
private static boolean notBlank(String value) {
|
||||
return value != null && !value.isBlank();
|
||||
}
|
||||
|
||||
@FunctionalInterface
|
||||
public interface RuntimeProvider {
|
||||
Optional<VisionRuntime> resolve();
|
||||
}
|
||||
|
||||
@FunctionalInterface
|
||||
public interface VisionCaller {
|
||||
VisionResponse send(VisionRuntime runtime, String requestBody) throws Exception;
|
||||
}
|
||||
|
||||
public record VisionRuntime(String modelName, String baseUrl, String apiKey) {
|
||||
}
|
||||
|
||||
public record VisionResponse(int statusCode, String body) {
|
||||
}
|
||||
|
||||
public static class OcrUnavailableException extends RuntimeException {
|
||||
private final String code;
|
||||
|
||||
public OcrUnavailableException(String code) {
|
||||
super(code);
|
||||
this.code = code;
|
||||
}
|
||||
|
||||
public OcrUnavailableException(String code, Throwable cause) {
|
||||
super(code, cause);
|
||||
this.code = code;
|
||||
}
|
||||
|
||||
public String code() {
|
||||
return code;
|
||||
}
|
||||
}
|
||||
}
|
||||
+74
@@ -0,0 +1,74 @@
|
||||
package org.dromara.aihr.personal;
|
||||
|
||||
import com.fasterxml.jackson.databind.ObjectMapper;
|
||||
import org.dromara.aihr.personal.service.PersonalVisionOcrService;
|
||||
import org.dromara.aihr.personal.service.PersonalVisionOcrService.OcrUnavailableException;
|
||||
import org.dromara.aihr.personal.service.PersonalVisionOcrService.VisionResponse;
|
||||
import org.dromara.aihr.personal.service.PersonalVisionOcrService.VisionRuntime;
|
||||
import org.junit.jupiter.api.Tag;
|
||||
import org.junit.jupiter.api.Test;
|
||||
|
||||
import java.util.Optional;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.assertEquals;
|
||||
import static org.junit.jupiter.api.Assertions.assertThrows;
|
||||
import static org.junit.jupiter.api.Assertions.assertTrue;
|
||||
|
||||
@Tag("dev")
|
||||
class PersonalVisionOcrServiceTest {
|
||||
|
||||
private static final byte[] JPEG = new byte[] {1, 2, 3};
|
||||
|
||||
@Test
|
||||
void recognizesAndNormalizesPageText() {
|
||||
AtomicReference<String> requestBody = new AtomicReference<>();
|
||||
PersonalVisionOcrService service = PersonalVisionOcrService.forTest(
|
||||
new ObjectMapper(), true, true,
|
||||
() -> Optional.of(new VisionRuntime("vision-model", "https://vision.example/v1", "secret")),
|
||||
(runtime, body) -> {
|
||||
requestBody.set(body);
|
||||
return new VisionResponse(200,
|
||||
"{\"choices\":[{\"message\":{\"content\":\" 第一条 \\n\\n 第二条 \"}}]}");
|
||||
});
|
||||
|
||||
assertEquals("第一条\n第二条", service.recognize(JPEG, "image/jpeg", 3));
|
||||
assertTrue(requestBody.get().contains("data:image/jpeg;base64,AQID"));
|
||||
assertTrue(requestBody.get().contains("第3页"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void rejectsWhenCostGuardIsDisabled() {
|
||||
PersonalVisionOcrService service = PersonalVisionOcrService.forTest(
|
||||
new ObjectMapper(), false, true,
|
||||
() -> Optional.of(new VisionRuntime("vision-model", "https://vision.example/v1", null)),
|
||||
(runtime, body) -> new VisionResponse(200, "{}"));
|
||||
|
||||
OcrUnavailableException error = assertThrows(OcrUnavailableException.class,
|
||||
() -> service.recognize(JPEG, "image/jpeg", 1));
|
||||
assertEquals("PERSONAL_OCR_RUNTIME_DISABLED", error.code());
|
||||
}
|
||||
|
||||
@Test
|
||||
void rejectsWhenNoVisionOrChatRuntimeExists() {
|
||||
PersonalVisionOcrService service = PersonalVisionOcrService.forTest(
|
||||
new ObjectMapper(), true, true, Optional::empty,
|
||||
(runtime, body) -> new VisionResponse(200, "{}"));
|
||||
|
||||
OcrUnavailableException error = assertThrows(OcrUnavailableException.class,
|
||||
() -> service.recognize(JPEG, "image/jpeg", 1));
|
||||
assertEquals("PERSONAL_OCR_MODEL_UNAVAILABLE", error.code());
|
||||
}
|
||||
|
||||
@Test
|
||||
void mapsProviderFailureToControlledCode() {
|
||||
PersonalVisionOcrService service = PersonalVisionOcrService.forTest(
|
||||
new ObjectMapper(), true, true,
|
||||
() -> Optional.of(new VisionRuntime("vision-model", "https://vision.example/v1", null)),
|
||||
(runtime, body) -> new VisionResponse(503, "provider unavailable"));
|
||||
|
||||
OcrUnavailableException error = assertThrows(OcrUnavailableException.class,
|
||||
() -> service.recognize(JPEG, "image/jpeg", 1));
|
||||
assertEquals("PERSONAL_OCR_PROVIDER_FAILED", error.code());
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user