added dalle

This commit is contained in:
gohumble 2022-08-30 16:15:14 +02:00
parent 98bc70fc7c
commit 92f9ba1550
4 changed files with 316 additions and 18 deletions

14
internal/dalle/client.go Normal file
View file

@ -0,0 +1,14 @@
package dalle
import (
"context"
"io"
)
type Client interface {
Generate(ctx context.Context, prompt string) (*Task, error)
ListTasks(ctx context.Context, req *ListTasksRequest) (*ListTasksResponse, error)
GetTask(ctx context.Context, taskID string) (*Task, error)
Download(ctx context.Context, generationID string) (io.ReadCloser, error)
Share(ctx context.Context, generationID string) (string, error)
}

11
internal/dalle/dalle.go Normal file
View file

@ -0,0 +1,11 @@
package dalle
const (
StatusPending = "pending"
StatusRejected = "rejected"
StatusSucceeded = "succeeded"
TaskTypeText2Im = "text2im"
defaultBatchSize = 4
)

View file

@ -0,0 +1,259 @@
package dalle
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"time"
)
const (
libraryVersion = "1.0.0"
defaultUserAgent = "dalle/" + libraryVersion
baseURL = "https://labs.openai.com/api/labs"
defaultHTTPClientTimeout = 15 * time.Second
)
type option func(*HTTPClient) error
func WithHTTPClient(httpClient *http.Client) option {
return func(c *HTTPClient) error {
c.httpClient = httpClient
return nil
}
}
func WithUserAgent(userAgent string) option {
return func(c *HTTPClient) error {
c.userAgent = userAgent
return nil
}
}
type HTTPClient struct {
httpClient *http.Client
userAgent string
apiKey string
}
var _ Client = (*HTTPClient)(nil)
func NewHTTPClient(apiKey string, opts ...option) (*HTTPClient, error) {
c := &HTTPClient{
httpClient: &http.Client{Timeout: defaultHTTPClientTimeout},
userAgent: defaultUserAgent,
apiKey: apiKey,
}
for _, opt := range opts {
if err := opt(c); err != nil {
return nil, err
}
}
return c, nil
}
type Task struct {
Object string `json:"object"`
ID string `json:"id"`
Created int64 `json:"created"`
TaskType string `json:"task_type"`
Status string `json:"status"`
PromptID string `json:"prompt_id"`
Prompt Prompt `json:"prompt"`
Generations Generations `json:"generations"`
}
type Generations struct {
Data []GenerationData `json:"data"`
Object string `json:"object"`
}
type GenerationData struct {
Created int64 `json:"created"`
Generation Generation `json:"generation"`
GenerationType string `json:"generation_type"`
ID string `json:"id"`
}
type Generation struct {
ImagePath string `json:"image_path"`
}
type Prompt struct {
ID string `json:"id"`
Object string `json:"object"`
Created int64 `json:"created"`
PromptType string `json:"prompt_type"`
Prompt struct {
Caption string `json:"caption"`
} `json:"prompt"`
ParentGenerationID string `json:"parent_generation_id"`
}
type GenerateRequest struct {
Prompt GenerateRequestPrompt `json:"prompt"`
TaskType string `json:"task_type"`
}
type GenerateRequestPrompt struct {
BatchSize int32 `json:"batch_size"`
Caption string `json:"caption"`
}
func (c *HTTPClient) Generate(ctx context.Context, caption string) (*Task, error) {
task := &Task{}
req := &GenerateRequest{
Prompt: GenerateRequestPrompt{
BatchSize: defaultBatchSize,
Caption: caption,
},
TaskType: TaskTypeText2Im,
}
return task, c.request(ctx, "POST", "/tasks", nil, req, task)
}
type ListTasksResponse struct {
Object string `json:"object"`
Data []Task `json:"data"`
}
type ListTasksRequest struct {
Limit int32 `json:"limit"`
}
func (c *HTTPClient) ListTasks(ctx context.Context, req *ListTasksRequest) (*ListTasksResponse, error) {
res := &ListTasksResponse{}
url := "/tasks"
if req != nil {
if req.Limit != 0 {
url += fmt.Sprintf("?limit=%d", req.Limit)
}
}
return res, c.request(ctx, "GET", url, nil, nil, res)
}
func (c *HTTPClient) GetTask(ctx context.Context, taskID string) (*Task, error) {
task := &Task{}
return task, c.request(ctx, "GET", "/tasks/"+taskID, nil, nil, task)
}
func (c *HTTPClient) Download(ctx context.Context, generationID string) (io.ReadCloser, error) {
req, err := c.createRequest(ctx, "/generations/"+generationID+"/download", "GET", nil, nil)
if err != nil {
return nil, fmt.Errorf("creating request: %w", err)
}
resp, err := c.httpClient.Do(req)
if err != nil {
return nil, fmt.Errorf("performing request: %w", err)
}
return resp.Body, nil
}
// Share makes the generation public and returns the public url
func (c *HTTPClient) Share(ctx context.Context, generationID string) (string, error) {
res := &GenerationData{}
err := c.request(ctx, "POST", "/generations/"+generationID+"/share", nil, nil, res)
if err != nil {
return "", err
}
return res.Generation.ImagePath, nil
}
func (c *HTTPClient) createRequest(ctx context.Context, path, method string, values *url.Values, data interface{}) (*http.Request, error) {
url := baseURL + path
if values != nil {
url += "?" + values.Encode()
}
var body io.Reader
if data != nil {
b, err := json.Marshal(data)
if err != nil {
return nil, fmt.Errorf("parsing request data: %w", err)
}
body = bytes.NewReader(b)
}
req, err := http.NewRequestWithContext(ctx, method, url, body)
if err != nil {
return nil, fmt.Errorf("building request: %w", err)
}
req.Header.Add("Authorization", "Bearer "+c.apiKey)
req.Header.Add("Content-Type", "application/json")
req.Header.Set("User-Agent", c.userAgent)
req.Header.Set("Authority", "labs.openai.com")
req.Header.Set("Accept", "*/*")
req.Header.Set("Accept-Language", "en-US,en;q=0.9,de;q=0.8")
req.Header.Set("Cache-Control", "no-cache")
req.Header.Set("Content-Length", "0")
req.Header.Set("Cookie", "")
req.Header.Set("Dnt", "1")
req.Header.Set("Origin", "https://labs.openai.com")
req.Header.Set("Pragma", "no-cache")
req.Header.Set("Sec-Ch-Ua", "\"Chromium\";v=\"104\", \" Not A;Brand\";v=\"99\", \"Google Chrome\";v=\"104\"")
req.Header.Set("Sec-Ch-Ua-Mobile", "?0")
req.Header.Set("Sec-Ch-Ua-Platform", "\"macOS\"")
req.Header.Set("Sec-Fetch-Dest", "empty")
req.Header.Set("Sec-Fetch-Mode", "cors")
req.Header.Set("Sec-Fetch-Site", "same-origin")
req.Header.Set("User-Agent", "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/104.0.0.0 Safari/537.36")
return req, nil
}
func (c *HTTPClient) request(ctx context.Context, method, path string, values *url.Values, body interface{}, result interface{}) error {
req, err := c.createRequest(ctx, path, method, values, body)
if err != nil {
return fmt.Errorf("creating request: %w", err)
}
resp, err := c.httpClient.Do(req)
if err != nil {
return fmt.Errorf("performing request: %w", err)
}
defer resp.Body.Close()
respBody, _ := io.ReadAll(resp.Body)
// TODO: improve error handling...
if resp.StatusCode != http.StatusOK {
return Error{
Message: "unexpected non 200 status code",
StatusCode: resp.StatusCode,
Details: string(respBody),
}
}
if err = json.Unmarshal(respBody, result); err != nil {
return Error{
Message: err.Error(),
StatusCode: resp.StatusCode,
Details: string(respBody),
}
}
return nil
}
type Error struct {
Message string
StatusCode int
Details string
}
func (e Error) Error() string {
return fmt.Sprintf("dalle: %s (status: %d, details: %s)", e.Message, e.StatusCode, e.Details)
}

View file

@ -4,12 +4,13 @@ import (
"bytes"
"context"
"fmt"
"github.com/LightningTipBot/LightningTipBot/internal/dalle"
"github.com/LightningTipBot/LightningTipBot/internal/runtime"
"github.com/LightningTipBot/LightningTipBot/internal/telegram/intercept"
"github.com/dillonstreator/dalle"
log "github.com/sirupsen/logrus"
"github.com/skip2/go-qrcode"
tb "gopkg.in/lightningtipbot/telebot.v3"
"io"
"os"
"time"
)
@ -24,7 +25,7 @@ func (bot *TipBot) generateImages(ctx intercept.Context) (intercept.Context, err
return ctx, err
}
invoice, err := bot.createInvoiceWithEvent(ctx, me, 1, fmt.Sprintf("DALLE2 %s", GetUserStr(user.Telegram)), InvoiceCallbackGenerateDalle, "")
invoice.Payer = user
if err != nil {
return ctx, err
}
@ -44,20 +45,22 @@ func (bot *TipBot) generateImages(ctx intercept.Context) (intercept.Context, err
func (bot *TipBot) generateDalleImages(event Event) {
invoiceEvent := event.(*InvoiceEvent)
user := invoiceEvent.User
user := invoiceEvent.Payer
if user.Wallet == nil {
return
}
// create the client with the bearer token api key
dalleClient, err := dalle.NewHTTPClient("")
// handle err
if err != nil {
return
}
ctx, cancel := context.WithTimeout(context.Background(), time.Minute*5)
defer cancel()
// generate a task to create an image with a prompt
task, err := dalleClient.Generate(ctx, "monkey printing bitcoin with money printing machine, cyberpunk")
task, err := dalleClient.Generate(ctx, "dogs fighting for bitcoin on sunny island, van gogh style")
// handle err
// poll the task.ID until status is succeeded
@ -67,7 +70,6 @@ func (bot *TipBot) generateDalleImages(event Event) {
t, err = dalleClient.GetTask(ctx, task.ID)
// handle err
t.Status = dalle.StatusSucceeded
if t.Status == dalle.StatusSucceeded {
fmt.Println("task succeeded")
@ -78,19 +80,31 @@ func (bot *TipBot) generateDalleImages(event Event) {
fmt.Println("task still pending")
}
/*
// download the first generated image
for _, data := range t.Generations.Data {
reader, err := dalleClient.Download(ctx, data.ID)
if err != nil {
continue
}
bot.trySendMessage(user.Telegram, &tb.Photo{File: tb.File{FileReader: reader}, Caption: fmt.Sprintf("Result")})
}*/
reader, err := os.OpenFile("image", 0, os.ModePerm)
if err != nil {
panic(err)
// download the first generated image
for _, data := range t.Generations.Data {
reader, err := dalleClient.Download(ctx, data.ID)
if err != nil {
return
}
defer reader.Close()
file, err := os.Create("images/" + data.ID + ".png")
if err != nil {
return
}
defer file.Close()
_, err = io.Copy(file, reader)
if err != nil {
return
}
f, err := os.OpenFile("images/"+data.ID+".png", 0, os.ModePerm)
if err != nil {
return
}
bot.trySendMessage(invoiceEvent.Payer.Telegram, &tb.Photo{File: tb.File{FileReader: f}, Caption: fmt.Sprintf("Result")})
}
bot.trySendMessage(user.Telegram, &tb.Photo{File: tb.File{FileReader: reader}, Caption: fmt.Sprintf("Result")})
// handle err and close readCloser
}