mirror of
https://github.com/ChuckNorrison/LightningTipBot.git
synced 2026-08-13 12:33:14 +02:00
added dalle
This commit is contained in:
parent
98bc70fc7c
commit
92f9ba1550
4 changed files with 316 additions and 18 deletions
14
internal/dalle/client.go
Normal file
14
internal/dalle/client.go
Normal 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
11
internal/dalle/dalle.go
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
package dalle
|
||||
|
||||
const (
|
||||
StatusPending = "pending"
|
||||
StatusRejected = "rejected"
|
||||
StatusSucceeded = "succeeded"
|
||||
|
||||
TaskTypeText2Im = "text2im"
|
||||
|
||||
defaultBatchSize = 4
|
||||
)
|
||||
259
internal/dalle/httpclient.go
Normal file
259
internal/dalle/httpclient.go
Normal 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)
|
||||
}
|
||||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue