Add OllamaClient implementation of AiClient

This commit is contained in:
imoyb
2023-10-02 22:56:45 +08:00
committed by Mark Pollack
parent c15663137c
commit e68bdeb9a0
6 changed files with 491 additions and 0 deletions

View File

@@ -16,6 +16,7 @@
<module>spring-ai-core</module>
<module>spring-ai-openai</module>
<module>spring-ai-azure-openai</module>
<module>spring-ai-ollama</module>
<module>spring-ai-spring-boot-autoconfigure</module>
<module>spring-ai-spring-boot-starters/spring-ai-starter-openai</module>
<module>spring-ai-spring-boot-starters/spring-ai-starter-azure-openai</module>

View File

@@ -0,0 +1,9 @@
## Ollama
Ollama lets you ge tup an running with large language models locally
Refer to the official [README](https://github.com/jmorganca/ollama) to get started.
Note, installing `ollama run llama2` will download a 4GB docker image.
You can run the disabled test in `OllamaClientTests.java` to kick the tires.

42
spring-ai-ollama/pom.xml Normal file
View File

@@ -0,0 +1,42 @@
<?xml version="1.0" encoding="UTF-8"?>
<project xmlns="http://maven.apache.org/POM/4.0.0"
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/xsd/maven-4.0.0.xsd">
<modelVersion>4.0.0</modelVersion>
<parent>
<groupId>org.springframework.experimental.ai</groupId>
<artifactId>spring-ai</artifactId>
<version>0.7.0-SNAPSHOT</version>
</parent>
<artifactId>spring-ai-ollama</artifactId>
<packaging>jar</packaging>
<name>Spring AI Ollama</name>
<description>Ollama support</description>
<properties>
<maven.compiler.source>17</maven.compiler.source>
<maven.compiler.target>17</maven.compiler.target>
<project.build.sourceEncoding>UTF-8</project.build.sourceEncoding>
</properties>
<dependencies>
<dependency>
<groupId>org.springframework.experimental.ai</groupId>
<artifactId>spring-ai-core</artifactId>
<version>${project.parent.version}</version>
</dependency>
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-logging</artifactId>
</dependency>
<!-- test dependencies -->
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>

View File

@@ -0,0 +1,248 @@
package org.springframework.ai.ollama.client;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.databind.ObjectMapper;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.client.AiClient;
import org.springframework.ai.client.AiResponse;
import org.springframework.ai.client.Generation;
import org.springframework.ai.prompt.Prompt;
import org.springframework.util.CollectionUtils;
import java.io.BufferedReader;
import java.io.IOException;
import java.io.InputStream;
import java.io.InputStreamReader;
import java.net.URI;
import java.net.http.HttpClient;
import java.net.http.HttpRequest;
import java.net.http.HttpResponse;
import java.time.Duration;
import java.util.*;
import java.util.function.Consumer;
import java.util.stream.Collectors;
/**
* A client implementation for interacting with Ollama Service. This class acts as an
* interface between the application and the Ollama AI Service, handling request creation,
* communication, and response processing.
*
* @author nullptr
*/
public class OllamaClient implements AiClient {
/** Logger for logging the events and messages. */
private static final Logger log = LoggerFactory.getLogger(OllamaClient.class);
/** Mapper for JSON serialization and deserialization. */
private static final ObjectMapper jsonMapper = new ObjectMapper();
/** HTTP client for making asynchronous calls to the Ollama Service. */
private static final HttpClient httpClient = HttpClient.newBuilder().build();
/** Base URL of the Ollama Service. */
private final String baseUrl;
/** Name of the model to be used for the AI service. */
private final String model;
/** Optional callback to handle individual generation results. */
private Consumer<OllamaGenerateResult> simpleCallback;
/**
* Constructs an OllamaClient with the specified base URL and model.
* @param baseUrl Base URL of the Ollama Service.
* @param model Model specification for the AI service.
*/
public OllamaClient(String baseUrl, String model) {
this.baseUrl = baseUrl;
this.model = model;
}
/**
* Constructs an OllamaClient with the specified base URL, model, and a callback.
* @param baseUrl Base URL of the Ollama Service.
* @param model Model specification for the AI service.
* @param simpleCallback Callback to handle individual generation results.
*/
public OllamaClient(String baseUrl, String model, Consumer<OllamaGenerateResult> simpleCallback) {
this(baseUrl, model);
this.simpleCallback = simpleCallback;
}
@Override
public AiResponse generate(Prompt prompt) {
validatePrompt(prompt);
HttpRequest request = buildHttpRequest(prompt);
var response = sendRequest(request);
List<OllamaGenerateResult> results = readGenerateResults(response.body());
return getAiResponse(results);
}
/**
* Validates the provided prompt.
* @param prompt The prompt to validate.
*/
protected void validatePrompt(Prompt prompt) {
if (CollectionUtils.isEmpty(prompt.getMessages())) {
throw new RuntimeException("The prompt message cannot be empty.");
}
if (prompt.getMessages().size() > 1) {
log.warn("Only the first prompt message will be used; subsequent messages will be ignored.");
}
}
/**
* Constructs an HTTP request for the provided prompt.
* @param prompt The prompt for which the request needs to be built.
* @return The constructed HttpRequest.
*/
protected HttpRequest buildHttpRequest(Prompt prompt) {
String requestBody = getGenerateRequestBody(prompt.getMessages().get(0).getContent());
// remove the suffix '/' if necessary
String url = !this.baseUrl.endsWith("/") ? this.baseUrl : this.baseUrl.substring(0, this.baseUrl.length() - 1);
return HttpRequest.newBuilder()
.uri(URI.create("%s/api/generate".formatted(url)))
.POST(HttpRequest.BodyPublishers.ofString(requestBody))
.timeout(Duration.ofMinutes(5L))
.build();
}
/**
* Sends the constructed HttpRequest and retrieves the HttpResponse.
* @param request The HttpRequest to be sent.
* @return HttpResponse containing the response data.
*/
protected HttpResponse<InputStream> sendRequest(HttpRequest request) {
var response = httpClient.sendAsync(request, HttpResponse.BodyHandlers.ofInputStream()).join();
if (response.statusCode() != 200) {
throw new RuntimeException("Ollama call returned an unexpected status: " + response.statusCode());
}
return response;
}
/**
* Serializes the prompt into a request body for the Ollama API call.
* @param prompt The prompt to be serialized.
* @return Serialized request body as a String.
*/
private String getGenerateRequestBody(String prompt) {
var data = Map.of("model", model, "prompt", prompt);
try {
return jsonMapper.writeValueAsString(data);
}
catch (JsonProcessingException ex) {
throw new RuntimeException("Failed to serialize the prompt to JSON", ex);
}
}
/**
* Reads and processes the results from the InputStream provided by the Ollama
* Service.
* @param inputStream InputStream containing the results from the Ollama Service.
* @return List of OllamaGenerateResult.
*/
protected List<OllamaGenerateResult> readGenerateResults(InputStream inputStream) {
try (BufferedReader bufferedReader = new BufferedReader(new InputStreamReader(inputStream))) {
var results = new ArrayList<OllamaGenerateResult>();
String line;
while ((line = bufferedReader.readLine()) != null) {
processResponseLine(line, results);
}
return results;
}
catch (IOException e) {
throw new RuntimeException("Error parsing Ollama generation response.", e);
}
}
/**
* Processes a single line from the Ollama response.
* @param line The line to be processed.
* @param results List to which parsed results will be added.
*/
protected void processResponseLine(String line, List<OllamaGenerateResult> results) {
if (line.isBlank())
return;
log.debug("Received ollama generate response: {}", line);
OllamaGenerateResult result;
try {
result = jsonMapper.readValue(line, OllamaGenerateResult.class);
}
catch (IOException e) {
throw new RuntimeException("Error parsing response line from Ollama.", e);
}
if (result.getModel() == null || result.getDone() == null) {
throw new IllegalStateException("Received invalid data from Ollama. Model = " + result.getModel()
+ " , Done = " + result.getDone());
}
if (simpleCallback != null) {
simpleCallback.accept(result);
}
results.add(result);
}
/**
* Converts the list of OllamaGenerateResult into a structured AiResponse.
* @param results List of OllamaGenerateResult.
* @return Formulated AiResponse.
*/
protected AiResponse getAiResponse(List<OllamaGenerateResult> results) {
var ollamaResponse = results.stream()
.filter(Objects::nonNull)
.filter(it -> it.getResponse() != null && !it.getResponse().isBlank())
.filter(it -> it.getDone() != null)
.map(OllamaGenerateResult::getResponse)
.collect(Collectors.joining(""));
var generation = new Generation(ollamaResponse);
// TODO investigate mapping of additional metadata/runtime info to the response.
// Determine if should be top
// level map vs. nested map
return new AiResponse(Collections.singletonList(generation), Map.of("ollama-generate-results", results));
}
/**
* @return Model name for the AI service.
*/
public String getModel() {
return model;
}
/**
* @return Base URL of the Ollama Service.
*/
public String getBaseUrl() {
return baseUrl;
}
/**
* @return Callback that handles individual generation results.
*/
public Consumer<OllamaGenerateResult> getSimpleCallback() {
return simpleCallback;
}
/**
* Sets the callback that handles individual generation results.
* @param simpleCallback The callback to be set.
*/
public void setSimpleCallback(Consumer<OllamaGenerateResult> simpleCallback) {
this.simpleCallback = simpleCallback;
}
}

View File

@@ -0,0 +1,146 @@
package org.springframework.ai.ollama.client;
import com.fasterxml.jackson.annotation.JsonIgnoreProperties;
import com.fasterxml.jackson.annotation.JsonProperty;
import java.util.List;
/**
* Ollama generate a completion api response model
*
* @author nullptr
*/
@JsonIgnoreProperties(ignoreUnknown = true)
public class OllamaGenerateResult {
@JsonProperty("model")
private String model;
@JsonProperty("created_at")
private String createdAt;
@JsonProperty("response")
private String response;
@JsonProperty("done")
private Boolean done;
@JsonProperty("context")
private List<Long> context;
@JsonProperty("total_duration")
private Long totalDuration;
@JsonProperty("load_duration")
private Long loadDuration;
@JsonProperty("prompt_eval_count")
private Long promptEvalCount;
@JsonProperty("prompt_eval_duration")
private Long promptEvalDuration;
@JsonProperty("eval_count")
private Long evalCount;
@JsonProperty("eval_duration")
private Long evalDuration;
public String getModel() {
return model;
}
public void setModel(String model) {
this.model = model;
}
public String getCreatedAt() {
return createdAt;
}
public void setCreatedAt(String createdAt) {
this.createdAt = createdAt;
}
public String getResponse() {
return response;
}
public void setResponse(String response) {
this.response = response;
}
public Boolean getDone() {
return done;
}
public void setDone(Boolean done) {
this.done = done;
}
public List<Long> getContext() {
return context;
}
public void setContext(List<Long> context) {
this.context = context;
}
public Long getTotalDuration() {
return totalDuration;
}
public void setTotalDuration(Long totalDuration) {
this.totalDuration = totalDuration;
}
public Long getLoadDuration() {
return loadDuration;
}
public void setLoadDuration(Long loadDuration) {
this.loadDuration = loadDuration;
}
public Long getPromptEvalCount() {
return promptEvalCount;
}
public void setPromptEvalCount(Long promptEvalCount) {
this.promptEvalCount = promptEvalCount;
}
public Long getPromptEvalDuration() {
return promptEvalDuration;
}
public void setPromptEvalDuration(Long promptEvalDuration) {
this.promptEvalDuration = promptEvalDuration;
}
public Long getEvalCount() {
return evalCount;
}
public void setEvalCount(Long evalCount) {
this.evalCount = evalCount;
}
public Long getEvalDuration() {
return evalDuration;
}
public void setEvalDuration(Long evalDuration) {
this.evalDuration = evalDuration;
}
@Override
public String toString() {
return "OllamaGenerateResult{" + "model='" + model + '\'' + ", createdAt='" + createdAt + '\'' + ", response='"
+ response + '\'' + ", done='" + done + '\'' + ", context=" + context + ", totalDuration="
+ totalDuration + ", loadDuration=" + loadDuration + ", promptEvalCount=" + promptEvalCount
+ ", promptEvalDuration=" + promptEvalDuration + ", evalCount=" + evalCount + ", evalDuration="
+ evalDuration + '}';
}
}

View File

@@ -0,0 +1,45 @@
package org.springframework.ai.ollama.client;
import org.junit.jupiter.api.Assertions;
import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.Test;
import org.springframework.ai.client.AiResponse;
import org.springframework.ai.prompt.Prompt;
import org.springframework.util.CollectionUtils;
import java.util.function.Consumer;
public class OllamaClientTests {
@Test
@Disabled("For manual smoke testing only.")
public void smokeTest() {
OllamaClient ollama2 = getOllamaClient();
Prompt prompt = new Prompt("Hello");
AiResponse aiResponse = ollama2.generate(prompt);
Assertions.assertNotNull(aiResponse);
Assertions.assertFalse(CollectionUtils.isEmpty(aiResponse.getGenerations()));
Assertions.assertFalse(CollectionUtils.isEmpty(aiResponse.getProviderOutput()));
Assertions.assertNotNull(aiResponse.getProviderOutput().get("ollama-generate-results"));
Assertions.assertNotNull(aiResponse.getGeneration());
Assertions.assertNotNull(aiResponse.getGeneration().getText());
}
private static OllamaClient getOllamaClient() {
Consumer<OllamaGenerateResult> ollamaGenerateResultConsumer = it -> {
if (it.getDone()) {
System.out.println();
System.out.printf("Total duration: %dms%n", it.getTotalDuration() / 1000 / 1000);
System.out.printf("Prompt tokens: %d%n", it.getPromptEvalCount());
System.out.printf("Completion tokens: %d%n", it.getEvalCount());
}
else {
System.out.print(it.getResponse());
}
};
return new OllamaClient("http://127.0.0.1:11434", "llama2", ollamaGenerateResultConsumer);
}
}