// Copyright (c) 2015-present Jeevanandam M (jeeva@myjeeva.com), All rights reserved.
// resty source code and usage is governed by a MIT style
// license that can be found in the LICENSE file.
// SPDX-License-Identifier: MIT

package resty

import (
	"bytes"
	"context"
	"fmt"
	"io"
	"mime"
	"mime/multipart"
	"net/http"
	"net/textproto"
	"net/url"
	"path"
	"path/filepath"
	"reflect"
	"strings"
)

//‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾
// Request Middleware(s)
//_______________________________________________________________________

// MiddlewareRequestCreate prepares the HTTP request from the user-provided [Request] values.
// It performs the following operations:
//   - Parse the request URL with path params and query params
//   - Parse the request headers from client and request level
//   - Parse the request body based on the content type and body type
//   - Create the underlying [http.Request] object
//   - Add credentials such as Basic Auth and Token Auth into the request
//
// Returns an error if request preparation fails.
func MiddlewareRequestCreate(c *Client, r *Request) (err error) {
	if err = parseRequestURL(c, r); err != nil {
		return err
	}

	// no error returned
	parseRequestHeader(c, r)

	if err = parseRequestBody(c, r); err != nil {
		return err
	}

	// possible error from `http.NewRequestWithContext` is an invalid
	// HTTP method since URL-related errors get caught up in the `parseRequestURL`
	if err = createRawRequest(c, r); err != nil {
		return err
	}

	addCredentials(c, r)

	_ = r.generateCurlCommand()

	return nil
}

func parseRequestURL(c *Client, r *Request) error {
	if len(c.PathParams())+len(r.PathParams) > 0 {
		// GitHub #103 Path Params, #663 Raw Path Params
		for p, v := range c.PathParams() {
			if _, ok := r.PathParams[p]; ok {
				continue
			}
			r.PathParams[p] = v
		}

		var prev int
		buf := acquireBuffer()
		defer releaseBuffer(buf)
		// search for the next or first opened curly bracket
		for curr := strings.Index(r.URL, "{"); curr == 0 || curr > prev; curr = prev + strings.Index(r.URL[prev:], "{") {
			// write everything from the previous position up to the current
			if curr > prev {
				buf.WriteString(r.URL[prev:curr])
			}
			// search for the closed curly bracket from current position
			next := curr + strings.Index(r.URL[curr:], "}")
			// if not found, then write the remainder and exit
			if next < curr {
				buf.WriteString(r.URL[curr:])
				prev = len(r.URL)
				break
			}
			// special case for {}, without parameter's name
			if next == curr+1 {
				buf.WriteString("{}")
			} else {
				// check for the replacement
				key := r.URL[curr+1 : next]
				value, ok := r.PathParams[key]
				// keep the original string if the replacement not found
				if !ok {
					value = r.URL[curr : next+1]
				}
				buf.WriteString(value)
			}

			// set the previous position after the closed curly bracket
			prev = next + 1
			if prev >= len(r.URL) {
				break
			}
		}
		if buf.Len() > 0 {
			// write remainder
			if prev < len(r.URL) {
				buf.WriteString(r.URL[prev:])
			}
			r.URL = buf.String()
		}
	}

	// Parsing request URL
	reqURL, err := url.Parse(r.URL)
	if err != nil {
		return &invalidRequestError{Err: err}
	}

	// If [Request.URL] is a relative path, then the following
	// gets evaluated in the order
	//	1. [Client.LoadBalancer] is used to obtain the base URL if not nil
	//	2. [Client.BaseURL] is used to obtain the base URL
	//	3. Otherwise [Request.URL] is used as-is
	if !reqURL.IsAbs() {
		r.URL = reqURL.String()
		if len(r.URL) > 0 && r.URL[0] != '/' {
			r.URL = "/" + r.URL
		}

		if r.client.LoadBalancer() != nil {
			r.baseURL, err = r.client.LoadBalancer().NextWithContext(r.Context())
			if err != nil {
				return &invalidRequestError{Err: err}
			}
		}

		reqURL, err = url.Parse(r.baseURL + r.URL)
		if err != nil {
			return &invalidRequestError{Err: err}
		}
	}

	// GH #407 && #318
	if reqURL.Scheme == "" && len(c.Scheme()) > 0 {
		reqURL.Scheme = c.Scheme()
	}

	// Adding Query Param
	if len(c.QueryParams())+len(r.QueryParams) > 0 {
		for k, v := range c.QueryParams() {
			if _, ok := r.QueryParams[k]; ok {
				continue
			}
			r.QueryParams[k] = v[:]
		}

		// GitHub #123 Preserve query string order partially.
		// Since not feasible in `SetQuery*` resty methods, because
		// standard package `url.Encode(...)` sorts the query params
		// alphabetically
		if isStringEmpty(reqURL.RawQuery) {
			reqURL.RawQuery = r.QueryParams.Encode()
		} else {
			reqURL.RawQuery = reqURL.RawQuery + "&" + r.QueryParams.Encode()
		}
	}

	// GH#797 Unescape query parameters (non-standard - not recommended)
	if r.unescapeQueryParams && len(reqURL.RawQuery) > 0 {
		// at this point, all errors caught up in the above operations
		// so ignore the return error on query unescape; I realized
		// while writing the unit test
		unescapedQuery, _ := url.QueryUnescape(reqURL.RawQuery)
		reqURL.RawQuery = strings.ReplaceAll(unescapedQuery, " ", "+") // otherwise request becomes bad request
	}

	r.URL = reqURL.String()

	return nil
}

func parseRequestHeader(c *Client, r *Request) {
	for k, v := range c.Header() {
		if _, ok := r.Header[k]; ok {
			continue
		}
		r.Header[k] = v[:]
	}

	if !r.isHeaderExists(hdrUserAgentKey) {
		r.Header.Set(hdrUserAgentKey, hdrUserAgentValue)
	}

	if !r.isHeaderExists(hdrAcceptEncodingKey) {
		r.Header.Set(hdrAcceptEncodingKey, r.client.ContentDecompresserKeys())
	}
}

func parseRequestBody(c *Client, r *Request) error {
	if r.isMultiPart && !(r.Method == MethodPost || r.Method == MethodPut || r.Method == MethodPatch) {
		err := fmt.Errorf("resty: multipart is not allowed in HTTP verb: %v", r.Method)
		return &invalidRequestError{Err: err}
	}

	if r.isPayloadSupported() {
		switch {
		case r.isMultiPart: // Handling Multipart
			if err := handleMultipart(c, r); err != nil {
				return &invalidRequestError{Err: err}
			}
		case len(c.FormData()) > 0 || len(r.FormData) > 0: // Handling Form Data
			handleFormData(c, r)
		case r.Body != nil: // Handling Request body
			if err := handleRequestBody(c, r); err != nil {
				return &invalidRequestError{Err: err}
			}
		}
	} else {
		r.Body = nil // if the payload is not supported by HTTP verb, set explicit nil
	}

	return nil
}

func createRawRequest(c *Client, r *Request) (err error) {
	// init client trace if enabled
	r.initTraceIfEnabled()

	if r.bodyBuf == nil {
		if reader, ok := r.Body.(io.Reader); ok {
			r.RawRequest, err = http.NewRequestWithContext(r.Context(), r.Method, r.URL, reader)
		} else {
			r.RawRequest, err = http.NewRequestWithContext(r.Context(), r.Method, r.URL, nil)
		}
	} else {
		r.RawRequest, err = http.NewRequestWithContext(r.Context(), r.Method, r.URL, r.bodyBuf)
	}

	if err != nil {
		return &invalidRequestError{Err: err}
	}

	// get the context reference back from underlying RawRequest
	r.SetContext(r.RawRequest.Context())

	// Assign close connection option
	r.RawRequest.Close = r.IsCloseConnection

	// Add headers into http request
	r.RawRequest.Header = r.Header.Clone()

	// Add cookies from client instance into http request
	for _, cookie := range c.Cookies() {
		r.RawRequest.AddCookie(cookie)
	}

	// Add cookies from request instance into http request
	for _, cookie := range r.Cookies {
		r.RawRequest.AddCookie(cookie)
	}

	// Set given content length value into the request
	if r.isContentLengthSet {
		r.RawRequest.ContentLength = r.contentLength
	} else {
		r.contentLength = r.RawRequest.ContentLength
	}

	return
}

func addCredentials(c *Client, r *Request) error {
	credentialsAdded := false
	// Basic Auth
	if r.credentials != nil {
		credentialsAdded = true
		r.RawRequest.SetBasicAuth(r.credentials.Username, r.credentials.Password)
	}

	// Build the token Auth header
	if !isStringEmpty(r.AuthToken) {
		credentialsAdded = true
		r.RawRequest.Header.Set(r.HeaderAuthorizationKey, strings.TrimSpace(r.AuthScheme+" "+r.AuthToken))
	}

	if !c.IsDisableWarn() && credentialsAdded {
		if r.RawRequest.URL.Scheme == "http" {
			r.log.Warnf("Using sensitive credentials in HTTP mode is not secure. Use HTTPS")
		}
	}

	return nil
}

var multipartWriteField = func(w *multipart.Writer, name, value string) error {
	return w.WriteField(name, value)
}

var multipartWriteFormData = func(w *multipart.Writer, r *Request) error {
	for k, v := range r.FormData {
		for _, iv := range v {
			if err := multipartWriteField(w, k, iv); err != nil {
				return err
			}
		}
	}
	return nil
}

var multipartCreatePart = func(w *multipart.Writer, h textproto.MIMEHeader) (io.Writer, error) {
	return w.CreatePart(h)
}

var multipartSetBoundary = func(w *multipart.Writer, r *Request) error {
	if isStringEmpty(r.multipartBoundary) {
		return nil
	}
	return w.SetBoundary(r.multipartBoundary)
}

var multipartPipeWriterClose = func(w *io.PipeWriter) error {
	return w.Close()
}

func handleMultipartFormData(r *Request) error {
	r.bodyBuf = acquireBuffer()
	mw := multipart.NewWriter(r.bodyBuf)
	defer mw.Close()

	// set custom multipart boundary if exists
	if err := multipartSetBoundary(mw, r); err != nil {
		return err
	}

	r.Header.Set(hdrContentTypeKey, mw.FormDataContentType())

	return multipartWriteFormData(mw, r)
}

func handleMultipart(c *Client, r *Request) error {
	for k, v := range c.FormData() {
		if _, ok := r.FormData[k]; ok {
			continue
		}
		r.FormData[k] = v[:]
	}

	if len(r.multipartFields) == 0 {
		return handleMultipartFormData(r)
	}

	// pre-process multipart fields to catch possible errors
	for _, mf := range r.multipartFields {
		if mf.isValues() {
			continue
		}

		if err := mf.openFile(); err != nil {
			return err
		}

		if err := mf.detectContentType(); err != nil {
			return err
		}
	}

	// multipart streaming
	br, bw := io.Pipe()
	mw := multipart.NewWriter(bw)
	r.Body = br

	// set custom multipart boundary if exists
	if err := multipartSetBoundary(mw, r); err != nil {
		closeq(bw)
		return err
	}

	r.Header.Set(hdrContentTypeKey, mw.FormDataContentType())

	r.multipartErrChan = make(chan error, 1)
	go func() {
		defer close(r.multipartErrChan)
		defer func() {
			if err := mw.Close(); err != nil {
				r.multipartErrChan <- err
			}
			if err := multipartPipeWriterClose(bw); err != nil {
				r.multipartErrChan <- err
			}
		}()

		if err := multipartWriteFormData(mw, r); err != nil {
			r.multipartErrChan <- err
			return
		}

		ctx, cancel := context.WithCancel(r.Context())
		r.multipartCancelFunc = cancel
		for _, mf := range r.multipartFields {
			if mf.isValues() {
				for _, v := range mf.Values {
					if err := multipartWriteField(mw, mf.Name, v); err != nil {
						r.multipartErrChan <- err
						return
					}
				}
				continue
			}

			partWriter, err := multipartCreatePart(mw, mf.createHeader())
			if err != nil {
				r.multipartErrChan <- err
				return
			}

			partWriter = mf.wrapProgressCallbackIfPresent(partWriter)
			if len(mf.tempBuf) > 0 {
				if _, err = partWriter.Write(mf.tempBuf); err != nil {
					r.multipartErrChan <- err
					return
				}
			}

			reader := &gracefulStopReader{ctx: ctx, r: mf.Reader}
			if _, err = ioCopy(partWriter, reader); err != nil {
				r.multipartErrChan <- err
				return
			}
		}
	}()

	return nil
}

func handleFormData(c *Client, r *Request) {
	for k, v := range c.FormData() {
		if _, ok := r.FormData[k]; ok {
			continue
		}
		r.FormData[k] = v[:]
	}

	r.bodyBuf = acquireBuffer()
	r.bodyBuf.WriteString(r.FormData.Encode())
	r.Header.Set(hdrContentTypeKey, formContentType)
	r.isFormData = true
}

func handleRequestBody(c *Client, r *Request) error {
	contentType := strings.ToLower(r.Header.Get(hdrContentTypeKey))
	if isStringEmpty(contentType) {
		// it is highly recommended that the user provide a request content-type
		// so that we can minimize memory allocation and compute.
		contentType = detectContentType(r.Body)
	}
	if !r.isHeaderExists(hdrContentTypeKey) {
		r.Header.Set(hdrContentTypeKey, contentType)
	}

	r.bodyBuf = acquireBuffer()

	switch body := r.Body.(type) {
	case io.Reader:
		// Resty v3 onwards io.Reader used as-is with the request body.
		releaseBuffer(r.bodyBuf)
		r.bodyBuf = nil

		// enable multiple reads if body is *bytes.Buffer
		if b, ok := r.Body.(*bytes.Buffer); ok {
			v := b.Bytes()
			r.Body = bytes.NewReader(v)
		}

		// do seek start for retry attempt if io.ReadSeeker
		// interface supported; otherwise fail loudly rather than
		// silently retrying the request with an empty body.
		if r.Attempt > 1 {
			rs, ok := r.Body.(io.ReadSeeker)
			if !ok {
				return ErrReaderNotSeekable
			}
			if _, err := rs.Seek(0, io.SeekStart); err != nil {
				return err
			}
		}
		return nil
	case []byte:
		r.bodyBuf.Write(body)
	case string:
		r.bodyBuf.Write([]byte(body))
	default:
		encKey := inferContentTypeMapKey(contentType)
		if jsonKey == encKey {
			if !r.jsonEscapeHTML {
				return encodeJSONEscapeHTML(r.bodyBuf, r.Body, r.jsonEscapeHTML)
			}
		} else if xmlKey == encKey {
			if inferKind(r.Body) != reflect.Struct {
				releaseBuffer(r.bodyBuf)
				r.bodyBuf = nil
				return ErrUnsupportedRequestBodyKind
			}
		}

		// user registered encoders with resty fallback key
		encFunc, found := c.inferContentTypeEncoder(contentType, encKey)
		if !found {
			releaseBuffer(r.bodyBuf)
			r.bodyBuf = nil
			return fmt.Errorf("resty: content-type encoder not found for %s", contentType)
		}
		if err := encFunc(r.bodyBuf, r.Body); err != nil {
			releaseBuffer(r.bodyBuf)
			r.bodyBuf = nil
			return err
		}
	}

	return nil
}

//‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾‾
// Response Middleware(s)
//_______________________________________________________________________

// MiddlewareResponseAutoParse parses the response body automatically using the
// Content-Type decoder registered via [Client.AddContentTypeDecoder].
// When [Request.SetResult], [Request.SetResultError], or [Client.SetResultError]
// is used, the body is automatically unmarshalled into the provided object.
func MiddlewareResponseAutoParse(c *Client, res *Response) (err error) {
	if (res.CascadeError != nil && (res.Request.isMultiPart && res.StatusCode() == 0)) ||
		res.Request.IsResponseDoNotParse {
		return // move on
	}

	if res.StatusCode() == http.StatusNoContent {
		res.Request.ResultError = nil
		return
	}

	rct := strings.ToLower(firstNonEmpty(
		res.Request.ResponseForceContentType,
		res.Header().Get(hdrContentTypeKey),
		res.Request.ResponseExpectContentType,
	))
	decKey := inferContentTypeMapKey(rct)
	decFunc, found := c.inferContentTypeDecoder(rct, decKey)
	if !found {
		// the Content-Type decoder is not found; just read all the body bytes
		err = res.readAll()
		return
	}

	// HTTP status code > 199 and < 300, considered as Result
	if res.IsStatusSuccess() && res.Request.Result != nil {
		res.Request.ResultError = nil
		defer closeq(res.Body)
		err = decFunc(res.Body, res.Request.Result)
		res.IsRead = true
		return
	}

	// HTTP status code > 399, considered as Error
	if res.IsStatusFailure() {
		// global error type registered at client-instance
		if res.Request.ResultError == nil {
			res.Request.ResultError = c.newErrorInterface()
		}

		if res.Request.ResultError != nil {
			defer closeq(res.Body)
			err = decFunc(res.Body, res.Request.ResultError)
			res.IsRead = true
			return
		}
	}

	return
}

var hostnameReplacer = strings.NewReplacer(":", "_", ".", "_")

func sanitizeResponseSaveFileNameFromHeader(file string) (string, error) {
	file = strings.TrimSpace(file)
	if isStringEmpty(file) {
		return "", nil
	}

	normalized := strings.ReplaceAll(file, "\\", "/")
	if strings.HasPrefix(normalized, "/") || isWindowsAbsPath(normalized) {
		return "", fmt.Errorf("resty: invalid Content-Disposition filename: absolute path is not allowed")
	}
	for _, s := range strings.Split(normalized, "/") {
		if s == ".." {
			return "", fmt.Errorf("resty: invalid Content-Disposition filename: parent directory traversal is not allowed")
		}
	}

	base := path.Base(normalized)
	if isStringEmpty(base) || base == "." {
		return "", fmt.Errorf("resty: invalid Content-Disposition filename")
	}

	return base, nil
}

func isWindowsAbsPath(file string) bool {
	if len(file) < 3 || file[1] != ':' {
		return false
	}

	drive := file[0]
	if (drive < 'A' || drive > 'Z') && (drive < 'a' || drive > 'z') {
		return false
	}

	return file[2] == '/'
}

func isPathWithinBaseDirectory(baseDir, target string) bool {
	rel, err := filepath.Rel(baseDir, target)
	if err != nil {
		return false
	}
	if rel == "." {
		return true
	}
	if rel == ".." {
		return false
	}
	return !strings.HasPrefix(rel, ".."+string(filepath.Separator))
}

// MiddlewareResponseSaveToFile writes the HTTP response body to a file.
// The filename is determined in the following order:
//   - [Request.SetResponseSaveFileName]
//   - Content-Disposition header
//   - Request URL path using [path.Base]
//   - Request URL hostname if the path is empty or "/"
//
// Content-Disposition filename values are sanitized before use. Absolute paths
// and parent-directory traversal segments are rejected.
//
// If [Client.SetResponseSaveDirectory] is set and
// [Request.SetResponseSaveFileName] provides a relative path, the final path
// must remain within the configured response save directory after cleaning,
// otherwise this middleware returns an error.
func MiddlewareResponseSaveToFile(c *Client, res *Response) error {
	if res.CascadeError != nil || !res.Request.IsResponseSaveToFile {
		return nil
	}

	file := res.Request.ResponseSaveFileName
	if isStringEmpty(file) {
		cntDispositionValue := res.Header().Get(hdrContentDisposition)
		if len(cntDispositionValue) > 0 {
			if _, params, err := mime.ParseMediaType(cntDispositionValue); err == nil {
				file, err = sanitizeResponseSaveFileNameFromHeader(params["filename"])
				if err != nil {
					return err
				}
			}
		}
		if isStringEmpty(file) {
			rURL, _ := url.Parse(res.Request.URL)
			if isStringEmpty(rURL.Path) || rURL.Path == "/" {
				file = hostnameReplacer.Replace(rURL.Host)
			} else {
				file = path.Base(rURL.Path)
			}
		}
	}

	baseDirRaw := c.ResponseSaveDirectory()
	hasBaseDir := !isStringEmpty(strings.TrimSpace(baseDirRaw))
	baseDir := ""
	if hasBaseDir {
		baseDir = filepath.Clean(baseDirRaw)
	}
	constrainToBaseDir := hasBaseDir && !isStringEmpty(res.Request.ResponseSaveFileName) && !filepath.IsAbs(file)

	if hasBaseDir && !filepath.IsAbs(file) {
		file = filepath.Join(baseDir, file)
	}

	file = filepath.Clean(file)
	if constrainToBaseDir && !isPathWithinBaseDirectory(baseDir, file) {
		return fmt.Errorf("resty: invalid save file path outside response save directory")
	}

	if err := createDirectory(filepath.Dir(file)); err != nil {
		return err
	}

	outFile, err := createFile(file)
	if err != nil {
		return err
	}

	defer func() {
		closeq(outFile)
		closeq(res.Body)
	}()

	// io.Copy reads maximum 32kb size, it is perfect for large file download too
	res.size, err = ioCopy(outFile, res.Body)

	return err
}
